ó
    >:j7  ã                  óè   • S SK Jr  S SKrS SKrS SKrS SKJrJr  S SKrS SK	J
r
  S SKJ
s  Jr  S SKJrJr  S SKJr  S SKJr  SSKJr  SS	KJr   " S
 S\
R0                  \5      r          SS jrg)é    )ÚannotationsN)ÚAnyÚOptional)ÚBaseTunerLayerÚcheck_adapters_to_merge)Ú	transpose)Úgather_params_ctxé   )Ú
LoraConfig)Ú	LoraLayerc                  óÐ   ^ • \ rS rSrSr   S         SU 4S jjjr\R                  SSS4               SS jjrSS jr	SSS jjr
SS	 jrSS
 jrSU 4S jjrSrU =r$ )ÚLoraParallelLinearé!   a|  
When the target layer parallel_linear is RowParallelLinear, in order to keep the input and output shapes
consistent, we need to split the lora matrix A into rows, and the lora_B at this time should be a complete linear
layer; In the same way, when the target layer is ColumnParallelLinear, we perform column segmentation on lora_B,
while lora_A is still a complete linear layer.
Fc           	     óê  >• UR                   (       a"  [        U R                  R                   S35      e[        TU ]  5         [        R
                  " U 4SU0UD6  UR                  (       a"  [        U R                  R                   S35      eX@l        [        XR                  5      U l        UR                  U l        X l        US   n	SU	0n
[        R                  n[!        U	S5      (       a  U	R"                  nSnSnU R                  (       a  UR$                  nOUR&                  nU R(                  " UU4UUUUUS.U
D6  U(       a"  [        U R                  R                   S	35      eSU l        g )
Nz0 does not support lora_bias yet, set it to FalseÚ
base_layerz2 does not support DoRA yet, please set it to FalseÚmegatron_configÚinit_methodTF)Ú
lora_alphaÚconfigr   Úinput_is_parallelÚgather_outputzB does not support target_conv_1d_layer yet, please set it to False)Ú	lora_biasÚ
ValueErrorÚ	__class__Ú__name__ÚsuperÚ__init__r   Úuse_doraÚbackendÚ
isinstanceÚRowParallelLinearÚis_parallel_aÚfan_in_fan_outÚ_active_adapterÚinitÚxavier_normal_Úhasattrr   r   r   Úupdate_layerÚis_target_conv_1d_layer)Úselfr   Úadapter_namer   r   Úrr   r)   Úkwargsr   Úparallel_linear_kwargsr   r   r   r   s                 €ÚV/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/tuners/lora/tp_layer.pyr   ÚLoraParallelLinear.__init__)   sj  ø€ ð ××Ü §¡× 7Ñ 7Ð8Ð8hÐiÓjÐjä‰ÑÔÜ×Ò˜4ÑA¨JÐA¸&ÒAà�?�?Ü §¡× 7Ñ 7Ð8Ð8jÐkÓlÐlàŒÜ'¨
×4MÑ4MÓNˆÔØ$×3Ñ3ˆÔØ+Ôà Ð!2Ñ3ˆØ"3°_Ð!EÐÜ×)Ñ)ˆÜ�? M×2Ñ2Ø)×5Ñ5ˆKØ ÐØˆØ××Ø *× <Ñ <Ñà&×4Ñ4ˆMØ×ÒØØð		
ð "ØØ#Ø/Ø'ñ		
ð %ò		
ö #ÜØ—>‘>×*Ñ*Ð+Ð+mÐnóð ð (-ˆÕ$ó    Tc	           
     ó,  • UR                   n
UR                  nUR                  nUR                  nUS::  a  [	        SU 35      eX R
                  U'   X0R                  U'   U
S:”  a  [        R                  " U
S9nO[        R                  " 5       nXàR                   U'   U	S   n[        R                  Ul        U R                  (       aX  U R                  R                  U R                   USUSUUS9n[        R"                  " X R$                  S[        R                  S	9nOW[        R"                  " U R                   US[        R                  S	9nU R                  R'                  UU R$                  SUUUS
9nUU R(                  U'   UU R*                  U'   U(       a'  U[,        R.                  " U5      -  U R0                  U'   OX2-  U R0                  U'   XÐR                  U'   [3        U[4        5      (       aU  UR7                  S5      (       a?  [9        U R;                  5       R<                  5         U R?                  X5        S S S 5        GO,[3        U[4        5      (       aT  UR7                  S5      (       a>  [9        U R;                  5       R<                  5         U RA                  X5        S S S 5        OÃ[3        U[4        5      (       aR  URC                  5       S:X  a>  [9        U R;                  5       R<                  5         U RE                  U5        S S S 5        O\US:X  a>  [9        U R;                  5       R<                  5         U RG                  U5        S S S 5        OU(       a  U RI                  X5        U RK                  U5        XRL                  ;   a  U RL                  U   RO                  XUS9  U RQ                  U RR                  US9  g ! , (       d  f       Nf= f! , (       d  f       Nw= f! , (       d  f       Nˆ= f! , (       d  f       N™= f)Nr   z?`r` should be a positive integer value but the value passed is g        )Úpr   FT)Ú
input_sizeÚoutput_sizeÚbiasr   Úskip_bias_addr   r   )Úin_featuresÚout_featuresr6   Údtype)r4   r5   r6   r   r   r   ÚpissaÚcordaÚoloraÚloftq)r+   r   )Úinference_mode)*Úlora_dropoutÚinit_lora_weightsÚ
use_rslorar   r   r,   r   ÚnnÚDropoutÚIdentityÚtorchÚfloat32Úparams_dtyper"   r   r!   r8   ÚLinearr9   ÚColumnParallelLinearÚlora_AÚlora_BÚmathÚsqrtÚscalingr    ÚstrÚ
startswithr	   Úget_base_layerÚweightÚ
pissa_initÚ
corda_initÚlowerÚ
olora_initÚ
loftq_initÚreset_lora_parametersÚ%_move_adapter_to_device_of_base_layerÚlora_variantr%   Úset_adapterÚactive_adapters)r*   r+   r,   r   r   r   r   r   r?   r.   r@   rA   rB   r   Úlora_dropout_layerr   Úlora_aÚlora_bs                     r/   r(   ÚLoraParallelLinear.update_layer^   s[  € ð ×*Ñ*ˆØ"×4Ñ4ÐØ×&Ñ&ˆ
Ø—?‘?ˆà�‹6ÜÐ^Ð_`Ð^aÐbÓcÐcØ �‰ˆ|ÑØ(2�‰˜Ñ%Ø˜#ÓÜ!#§¢¨lÑ!;Ñä!#§¢£Ðà*<×Ñ˜,Ñ'à0Ð1BÑCˆä',§}¡}ˆÔ$Ø××Ø—\‘\×3Ñ3Ø×+Ñ+ØØØ"3Ø"Ø'Ø&ð 4ð ˆFô —Y’Y¨1×;LÑ;LÐSXÔ`e×`mÑ`mÑn‰Fä—Y’Y¨4×+;Ñ+;È!ÐRWÔ_d×_lÑ_lÑmˆFØ—\‘\×6Ñ6ØØ ×-Ñ-ØØ+Ø'Ø&ð 7ð ˆFð %+ˆ�‰�LÑ!Ø$*ˆ�‰�LÑ!ÞØ)3´d·i²iÀ³lÑ)BˆD�L‰L˜Ò&à)3©ˆD�L‰L˜Ñ&à&.�‰�lÑ#ô Ð'¬×-Ñ-Ð2C×2NÑ2NÈw×2WÑ2WÜ" 4×#6Ñ#6Ó#8×#?Ñ#?Õ@Ø—‘ Ô@÷ AÑ@äÐ)¬3×/Ñ/Ð4E×4PÑ4PÐQX×4YÑ4YÜ" 4×#6Ñ#6Ó#8×#?Ñ#?Õ@Ø—‘ Ô@÷ AÐ@äÐ)¬3×/Ñ/Ð4E×4KÑ4KÓ4MÐQXÓ4XÜ" 4×#6Ñ#6Ó#8×#?Ñ#?Õ@Ø—‘ Ô-÷ AÐ@à 'Ó)Ü" 4×#6Ñ#6Ó#8×#?Ñ#?Õ@Ø—‘ Ô-÷ AÐ@æØ×&Ñ& |ÔGð 	×2Ñ2°<Ô@à×,Ñ,Ó,Ø×Ñ˜lÑ+×0Ñ0°ÐY_Ð0Ñ`à×Ñ˜×-Ñ-¸nÐÒM÷) AÕ@ú÷ AÕ@ú÷ AÕ@ú÷ AÕ@ús0   È2OÊO#ÌO4ÍPÏ
O Ï#
O1Ï4
PÐ
Pc           	     óV  • U R                   " U/UQ70 UD6  UR                  SS 5      nU R                  (       a<  U R                  (       a  U R	                  5         U R
                  " U/UQ70 UD6u  pVXV4$ Ub"  [        U R                  R                   S35      eU R                  (       a  U R
                  " U/UQ70 UD6u  pVXV4$ U R
                  " U/UQ70 UD6u  pVUR                  nU R                   Hœ  nX€R                  R                  5       ;  a  M"  U R                  U   n	U R                  U   n
U R                  U   nU R                  U   nU R!                  XR"                  R                  5      nXZ" U	" U" U5      5      5      U-  -   nMž     UR%                  U5      nXV4$ )NÚadapter_namesz* does not support mixed_batch_forward yet.)Ú_check_forward_argsÚpopÚdisable_adaptersÚmergedÚunmerger   r   r   r   r:   r]   rK   ÚkeysrL   r@   rO   Ú_cast_input_dtyperS   Úto)r*   ÚxÚargsr-   rc   Úresultr6   Útorch_result_dtypeÚactive_adapterrK   rL   ÚdropoutrO   s                r/   ÚforwardÚLoraParallelLinear.forward³   s�  € Ø× Ò  Ð4 TÒ4¨VÒ4ØŸ
™
 ?°DÓ9ˆð × × Ø�{�{Ø—‘”ØŸ?š?¨1Ð>¨tÒ>°vÑ>‰LˆFð& ˆ|Ðð% Ñ&Ü §¡× 7Ñ 7Ð8Ð8bÐcÓdÐdØ�[�[ØŸ?š?¨1Ð>¨tÒ>°vÑ>‰LˆFð ˆ|Ðð  Ÿ?š?¨1Ð>¨tÒ>°vÑ>‰LˆFØ!'§¡ÐØ"&×"6Ô"6�Ø!¯©×)9Ñ)9Ó);Ó;ÙØŸ™ ^Ñ4�ØŸ™ ^Ñ4�Ø×+Ñ+¨NÑ;�ØŸ,™, ~Ñ6�Ø×*Ñ*¨1¯m©m×.AÑ.AÓB�Ø &©±¸³
Ó);Ó"<¸wÑ"FÑF’ñ #7ð —Y‘YÐ1Ó2ˆFØˆ|Ðr1   c                óX  • [        X5      nU(       d  gU GH  nX0R                  R                  5       ;   d  M#  U R                  5       nU(       a‚  UR                  R
                  R                  5       nU R                  U5      nXV-   n[        R                  " U5      R                  5       (       d  [        SU S35      eXTR                  l        O9U R                  U5      nUR                  R
                  U-   UR                  l        U R                  R                  U5        GM     g)a  
Merge the active adapter weights into the base weights

Args:
    safe_merge (`bool`, *optional*):
        If True, the merge operation will be performed in a copy of the original weights and check for NaNs
        before merging the weights. This is useful if you want to check if the merge operation will produce
        NaNs. Defaults to `False`.
    adapter_names (`list[str]`, *optional*):
        The list of adapter names that should be merged. If None, all active adapters will be merged. Defaults
        to `None`.
Nz1NaNs detected in the merged weights. The adapter z seems to be broken)r   rK   ri   rR   rS   ÚdataÚcloneÚget_delta_weightrF   ÚisfiniteÚallr   Úmerged_adaptersÚappend)r*   Ú
safe_mergerc   rp   r   Úorig_weightsÚdelta_weights          r/   ÚmergeÚLoraParallelLinear.mergeÑ   sû   € ô 0°ÓDˆÞàä+ˆNØ§¡×!1Ñ!1Ó!3Õ3Ø!×0Ñ0Ó2�
Þð $.×#4Ñ#4×#9Ñ#9×#?Ñ#?Ó#A�LØ#'×#8Ñ#8¸Ó#H�LØ#/Ñ#>�Lä Ÿ>š>¨,Ó7×;Ñ;×=Ñ=Ü(ØOÐP^ÐO_Ð_rÐsóð ð .:×%Ñ%Õ*à#'×#8Ñ#8¸Ó#H�LØ-7×->Ñ->×-CÑ-CÀlÑ-R�J×%Ñ%Ô*à×$Ñ$×+Ñ+¨N×;ò) ,r1   c                ó¬  • U R                   (       d  [        R                  " S5        g[        U R                  5      S:”  a“  U R                  R                  5       nXR                  R                  5       ;   a@  U R                  5       R                  nU R                  U5      nU=R                  U-  sl        [        U R                  5      S:”  a  M’  gg)zG
This method unmerges all merged adapter layers from the base weights.
z Already unmerged. Nothing to do.Nr   )rg   ÚwarningsÚwarnÚlenrz   re   rK   ri   rR   rS   rw   ru   )r*   rp   rS   r~   s       r/   rh   ÚLoraParallelLinear.unmergeù   sœ   € ð �{�{Ü�MŠMÐ<Ô=ØÜ�$×&Ñ&Ó'¨!Ó+Ø!×1Ñ1×5Ñ5Ó7ˆNØ§¡×!1Ñ!1Ó!3Ó3Ø×,Ñ,Ó.×5Ñ5�Ø#×4Ñ4°^ÓD�Ø—’˜|Ñ+•ô �$×&Ñ&Ó'¨!×+r1   c                óú  • U R                   U   R                  R                  nU R                   U   R                  R                  nUR                  S:H  =(       a-    U[
        R                  :H  =(       d    U[
        R                  :H  nU R                  U   R                  nU R                   U   R                  nU(       a   UR                  5       nUR                  5       n[        Xe-  U R                  5      U R                  U   -  nU(       ai  UR                  US9nUR                  U5      U R                  U   R                  l        UR                  U5      U R                   U   R                  l        U$ )zš
Compute the delta weight for the given adapter.

Args:
    adapter (str):
        The name of the adapter for which the delta weight should be computed.
Úcpu)r:   )rL   rS   Údevicer:   ÚtyperF   Úfloat16Úbfloat16rK   Úfloatr   r#   rO   rk   ru   )r*   Úadapterrˆ   r:   Úcast_to_fp32Úweight_AÚweight_BÚoutput_tensors           r/   rw   Ú#LoraParallelLinear.get_delta_weight  s,  € ð —‘˜WÑ%×,Ñ,×3Ñ3ˆØ—‘˜GÑ$×+Ñ+×1Ñ1ˆð
 —{‘{ eÑ+×c°¼%¿-¹-Ñ1G×1bÈ5ÔTY×TbÑTbÑKbˆà—;‘;˜wÑ'×.Ñ.ˆØ—;‘;˜wÑ'×.Ñ.ˆæØ—~‘~Ó'ˆHØ—~‘~Ó'ˆHä! (Ñ"5°t×7JÑ7JÓKÈdÏlÉlÐ[bÑNcÑcˆæØ)×,Ñ,°5Ð,Ð9ˆMð 08¯{©{¸5Ó/AˆD�K‰K˜Ñ ×'Ñ'Ô,Ø/7¯{©{¸5Ó/AˆD�K‰K˜Ñ ×'Ñ'Ô,àÐr1   c                ó*   >• [         TU ]  5       nSU-   $ )Nzlora.)r   Ú__repr__)r*   Úrepr   s     €r/   r”   ÚLoraParallelLinear.__repr__)  s   ø€ Ü‰gÑÓ ˆØ˜‰}Ðr1   )r$   r   r#   r"   r)   )r   r
   F)
r+   rP   r   r   r,   Úintr   r—   r)   Úbool)r+   rP   r,   r—   r   r—   r   r   r   r˜   r   r˜   r?   r˜   ÚreturnÚNone)rl   útorch.Tensorrm   r   r-   r   )FN)r|   r˜   rc   zOptional[list[str]]r™   rš   )r™   rš   )r™   r›   )r™   rP   )r   Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   r%   r&   r(   rr   r   rh   rw   r”   Ú__static_attributes__Ú__classcell__)r   s   @r/   r   r   !   sê   ø† ñð ØØ(-ð3-ð ð3-ð ð	3-ð ð3-ð ð3-ð "&÷3-ð 3-ðv ×'Ñ'Ø"&Ø#Ø$ðSNàðSNð ðSNð ð	SNð
 ðSNð  ðSNð ðSNð ðSNð 
õSNôjö<&<ôP,ô ÷Dõ r1   r   c                ó   • S n[        U [        5      (       a  U R                  5       nOU nUR                  (       a!  [        R
                  " UR                  5      nOS nU(       aè  [        UUR                  R                  UR                  R                  45      (       a­  UR                  5       nUR                  n[        U[        5      (       a2  UR                  R                  R                  n	U	" S0 UR                  D6nX‡S'   US   (       a"  [        R                   " S5        S=US'   Ul        [%        SU UUUR                  S.UD6nU$ )Nr   r#   z†fan_in_fan_out is set to True but the target module is `ColumnParallelLinear` or `RowParallelLinear`. Setting fan_in_fan_out to False.F)r   r+   r   r   © )r    r   rR   r   Ú	importlibÚimport_moduleÚmegatron_coreÚtensor_parallelrJ   r!   ÚcopyÚdictÚtransformerÚtransformer_configÚTransformerConfigr‚   rƒ   r#   r   )
Útargetr+   r   r-   Ú
new_moduleÚtarget_base_layerr¦   Úmegatron_kwargsr   Útransformer_config_classs
             r/   Údispatch_megatronr²   .  s>  € ð €Jä�&œ.×)Ñ)Ø"×1Ñ1Ó3Ñà"Ðà××Ü!×/Ò/°×0DÑ0DÓE‰àˆæœØØ	×	&Ñ	&×	;Ñ	;¸]×=ZÑ=Z×=lÑ=lÐm÷ñ ð !Ÿ+™+›-ˆØ ×0Ñ0ˆÜ�o¤t×,Ñ,Ø'4×'@Ñ'@×'SÑ'S×'eÑ'eÐ$Ù6ÑP¸×9OÑ9OÑPˆOØ-<Ð)Ñ*ØÐ+×,Ü�MŠMð3ôð
 INÐMˆOÐ,Ñ-°Ô0EÜ'ð 
ØØ%ØØ!×1Ñ1ñ	
ð
 ñ
ˆ
ð Ðr1   )
r­   ztorch.nn.Moduler+   rP   r   r   r-   r   r™   zOptional[torch.nn.Module])Ú
__future__r   r¤   rM   r‚   Útypingr   r   rF   Útorch.nnrC   Útorch.nn.initr%   Úpeft.tuners.tuners_utilsr   r   Ú
peft.utilsr   Úpeft.utils.integrationsr	   r   r   Úlayerr   ÚModuler   r²   r£   r1   r/   Ú<module>r¼      sy   ðõ #ã Û Û ß  ã Ý ß Ð ç LÝ  Ý 5å Ý ôJ˜Ÿ™ Iô JðZ+Øð+àð+ð ð+ð ð	+ð
 õ+r1   