ó
    >:j¥&  ã                   ó|  • S SK r S SKJr  S SKrS rS\R
                  S\S\R
                  4S jrS\R
                  S\S\S\R
                  4S	 jr	 SS\R
                  S\S
\S   S\S\R
                  4
S jjr
 SS\R
                  S
\S   S\R
                  4S jjrS\R
                  S\R
                  S\R
                  4S jrS\\R
                     S\R
                  S\R
                  4S jrS\\R
                     S\R
                  S\S\R
                  4S jr SS\\R
                     S\R
                  S\S\S   S\R
                  4
S jjrS\\R
                     S\R
                  S\S\R
                  4S jr SS\\R
                     S\R
                  S\S\S   S\R
                  4
S jjrg)é    N)ÚLiteralc                 óŠ   • UR                   SU R                  5       UR                  5       -
  -  -   nUR                  U5      nU$ )a-  
Reshapes `weights` to match the shape of `task_tensors` by unsqeezing in the remaining dimenions.

Args:
    task_tensors (`torch.Tensor`): The tensors that will be used to reshape `weights`.
    weights (`torch.Tensor`): The tensor to be reshaped.

Returns:
    `torch.Tensor`: The reshaped tensor.
)é   )ÚshapeÚdimÚview)Útask_tensorsÚweightsÚ	new_shapes      ÚS/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/utils/merge_utils.pyÚreshape_weight_task_tensorsr      s>   € ð —‘ ¨×(8Ñ(8Ó(:¸W¿[¹[»]Ñ(JÑ KÑK€IØ�l‰l˜9Ó%€GØ€Nó    ÚtensorÚdensityÚreturnc                 ó0  • [         R                  " U 5      R                  S5      n[        XR	                  5       -  5      n[         R
                  " U R                  5       R                  S5      USS9nSX$S   '   XR                  U R                  5      -  $ )a>  
Prune the smallest values of the task tensors and retain the top-k values based on the specified fraction
`density`.

Args:
    tensor (`torch.Tensor`):The tensor to prune.
    density (`float`):The fraction of values to preserve. Should be in [0,1].

Returns:
    `torch.Tensor`: The tensor with the pruned weights.
éÿÿÿÿT)ÚkÚlargestr   )ÚtorchÚ
zeros_likeÚreshapeÚintÚnumelÚtopkÚabsr   )r   r   Úmaskr   Útop_ks        r   Úmagnitude_based_pruningr   %   sv   € ô ×Ò˜FÓ#×+Ñ+¨BÓ/€DÜˆG—l‘l“nÑ$Ó%€AÜ�JŠJ�v—z‘z“|×+Ñ+¨BÓ/°1¸dÑC€EØ€Dˆq‰�NØ—L‘L §¡Ó.Ñ.Ð.r   Úrescalec                 ót   • [         R                  " [         R                  " XS95      nX-  nU(       a  XA-  nU$ )aa  
Prune random values based on the specified fraction `density`.

Args:
    tensor (`torch.Tensor`):The tensor to prune.
    density (`float`):The fraction of values to preserve. Should be in [0,1].
    rescale (`bool`):Whether to rescale the result to preserve the expected value of the original tensor.

Returns:
    `torch.Tensor`: The pruned tensor.
)ÚinputÚ
fill_value)r   Ú	bernoulliÚ	full_like)r   r   r    r   Úpruned_tensors        r   Úrandom_pruningr'   8   s3   € ô �?Š?œ5Ÿ?š?°ÑLÓM€DØ‘M€MÞØ%Ñ/ˆØÐr   Úmethod)Ú	magnitudeÚrandomc                 óÌ   • US:¼  a  [         R                  " SU S35        U $ US:  a  [        SU 35      eUS:X  a  [        X5      $ US:X  a
  [	        XUS9$ [        S	U 35      e)
a³  
Prune the values of task tensors based on the `method`.

Args:
    tensor (`torch.Tensor`):The tensor to prune.
    density (`float`):The fraction of values to preserve. Should be in [0,1].
    method (`str`):The method to use to prune. Should be one of ["magnitude", "random"].
    rescale (`bool`):Whether to rescale the result to preserve the expected value of the original tensor.

Returns:
    `torch.Tensor`: The pruned tensor.
r   zThe density z= is greater than or equal to 1, no pruning will be performed.r   zDensity should be >= 0, got r)   r*   )r    zUnknown method )ÚwarningsÚwarnÚ
ValueErrorr   r'   )r   r   r(   r    s       r   Úpruner/   K   sz   € ð �!ƒ|Ü�Š˜ W IÐ-jÐkÔlØˆØ	�1‹ÜÐ7¸°yÐAÓBÐBØ�ÓÜ& vÓ7Ð7Ø	�8Ó	Ü˜f°wÑ?Ð?ä˜?¨6¨(Ð3Ó4Ð4r   )ÚtotalÚ	frequencyc                 óÖ   • U R                  5       nUS:X  a  U R                  SS9nO%US:X  a  UR                  SS9nO[        SU S35      e[        R                  " US:¬  SS5      nX$:H  $ )	a>  
Get the mask of the majority sign across the task tensors. Task tensors are stacked on dimension 0.

Args:
    tensor (`torch.Tensor`):The tensor to get the mask from.
    method (`str`):The method to use to get the mask. Should be one of ["total", "frequency"].

Returns:
    `torch.Tensor`: The majority sign mask.
r0   r   ©r   r1   zUnimplemented mask method "Ú"r   r   )ÚsignÚsumÚRuntimeErrorr   Úwhere)r   r(   r5   Úsign_magnitudeÚmajority_signs        r   Úcalculate_majority_sign_maskr;   g   ss   € ð �;‰;‹=€DØ�ÓØŸ™¨˜Ð*‰Ø	�;Ó	ØŸ™ a˜˜‰äÐ8¸¸ÀÐBÓCÐCÜ—K’K °!Ñ 3°Q¸Ó;€MØÑ Ð r   r	   Úmajority_sign_maskc                 ór   • X-  R                  SS9nUR                  SS9nU[        R                  " USS9-  $ )a  
Merge the task tensors using disjoint merge.

Args:
    task_tensors (`torch.Tensor`):The task tensors to merge.
    majority_sign_mask (`torch.Tensor`):The mask of the majority sign across the task tensors.

Returns:
    `torch.Tensor`: The merged tensor.
r   r3   g      ð?)Úmin)r6   r   Úclamp)r	   r<   Úmixed_task_tensorsÚnum_params_preserveds       r   Údisjoint_mergerB   €   sF   € ð 'Ñ;×@Ñ@ÀQÐ@ÐGÐØ-×1Ñ1°aÐ1Ð8ÐØ¤§¢Ð,@ÀcÑ JÑJÐJr   r
   c                 ól   • [         R                  " U SS9n [        X5      nX-  nUR                  SS9nU$ )zé
Merge the task tensors using `task arithmetic`.

Args:
    task_tensors(`List[torch.Tensor]`):The task tensors to merge.
    weights (`torch.Tensor`):The weights of the task tensors.

Returns:
    `torch.Tensor`: The merged tensor.
r   r3   )r   Ústackr   r6   )r	   r
   Úweighted_task_tensorsr@   s       r   Útask_arithmeticrF   �   sA   € ô —;’;˜|°Ñ3€Lä)¨,Ó@€GØ(Ñ2ÐØ.×2Ñ2°qÐ2Ð9ÐØÐr   c           	      óª   • U  Vs/ s H  n[        X2SS9PM     n n[        R                  " U SS9n [        X5      nX-  nUR	                  SS9nU$ s  snf )a8  
Merge the task tensors using `task arithmetic`.

Args:
    task_tensors(`List[torch.Tensor]`):The task tensors to merge.
    weights (`torch.Tensor`):The weights of the task tensors.
    density (`float`): The fraction of values to preserve. Should be in [0,1].

Returns:
    `torch.Tensor`: The merged tensor.
r)   ©r(   r   r3   ©r/   r   rD   r   r6   ©r	   r
   r   r   rE   r@   s         r   Úmagnitude_prunerK   £   sd   € ñ NZÓZÊ\À6”E˜&°+Ô>É\€LÐZÜ—;’;˜|°Ñ3€Lä)¨,Ó@€GØ(Ñ2ÐØ.×2Ñ2°qÐ2Ð9ÐØÐùò [s   …AÚmajority_sign_methodc           	      ó´   • U  Vs/ s H  n[        XBSS9PM     n n[        R                  " U SS9n [        XS9n[	        X5      nX-  n[        Xe5      nU$ s  snf )a°  
Merge the task tensors using `ties`.

Args:
    task_tensors(`List[torch.Tensor]`):The task tensors to merge.
    weights (`torch.Tensor`):The weights of the task tensors.
    density (`float`):The fraction of values to preserve. Should be in [0,1].
    majority_sign_method (`str`):
        The method to use to get the majority sign mask. Should be one of ["total", "frequency"].

Returns:
    `torch.Tensor`: The merged tensor.
r)   rH   r   r3   ©r/   r   rD   r;   r   rB   ©r	   r
   r   rL   r   r<   rE   r@   s           r   ÚtiesrP   ¹   sg   € ñ( NZÓZÊ\À6”E˜&°+Ô>É\€LÐZÜ—;’;˜|°Ñ3€Lä5°lÑ`Ðä)¨,Ó@€GØ(Ñ2Ðä'Ð(=ÓRÐØÐùò [s   …Ac           
      ó¬   • U  Vs/ s H  n[        X2SSS9PM     n n[        R                  " U SS9n [        X5      nX-  nUR	                  SS9nU$ s  snf )a3  
Merge the task tensors using `dare linear`.

Args:
    task_tensors(`List[torch.Tensor]`):The task tensors to merge.
    weights (`torch.Tensor`):The weights of the task tensors.
    density (`float`):The fraction of values to preserve. Should be in [0,1].

Returns:
    `torch.Tensor`: The merged tensor.
r*   T©r(   r    r   r3   rI   rJ   s         r   Údare_linearrS   Ù   sh   € ñ YeÓeÒXdÈf”E˜&°(ÀDÔIÑXd€LÐeÜ—;’;˜|°Ñ3€Lä)¨,Ó@€GØ(Ñ2ÐØ.×2Ñ2°qÐ2Ð9ÐØÐùò fs   …Ac           
      ó¶   • U  Vs/ s H  n[        XBSSS9PM     n n[        R                  " U SS9n [        XS9n[	        X5      nX-  n[        Xe5      nU$ s  snf )aµ  
Merge the task tensors using `dare ties`.

Args:
    task_tensors(`List[torch.Tensor]`):The task tensors to merge.
    weights (`torch.Tensor`):The weights of the task tensors.
    density (`float`):The fraction of values to preserve. Should be in [0,1].
    majority_sign_method (`str`):
        The method to use to get the majority sign mask. Should be one of ["total", "frequency"].

Returns:
    `torch.Tensor`: The merged tensor.
r*   TrR   r   r3   rH   rN   rO   s           r   Ú	dare_tiesrU   ï   sk   € ñ( YeÓeÒXdÈf”E˜&°(ÀDÔIÑXd€LÐeÜ—;’;˜|°Ñ3€Lä5°lÑ`Ðä)¨,Ó@€GØ(Ñ2Ðä'Ð(=ÓRÐØÐùò fs   …A)F)r0   )r,   Útypingr   r   r   ÚTensorÚfloatr   Úboolr'   r/   r;   rB   ÚlistrF   rK   rP   rS   rU   © r   r   Ú<module>r\      s<  ðó Ý ã òð / E§L¡Lð /¸5ð /ÀUÇ\Á\ô /ð&˜5Ÿ<™<ð °%ð À$ð È5Ï<É<ô ð( chñ5Ø�L‰Lð5Ø#(ð5Ø29Ð:OÑ2Pð5Ø[_ð5à
‡\�\õ5ð: CJñ!Ø�L‰Lð!Ø")Ð*>Ñ"?ð!à
‡\�\õ!ð2K §¡ð KÀ5Ç<Á<ð KÐTY×T`ÑT`ô Kð  $ u§|¡|Ñ"4ð ¸u¿|¹|ð ÐPU×P\ÑP\ô ð& $ u§|¡|Ñ"4ð ¸u¿|¹|ð ÐV[ð Ð`e×`lÑ`lô ð4 ;Bñ	Ø�u—|‘|Ñ$ðà�\‰\ðð ðð "Ð"6Ñ7ð	ð
 ‡\�\õð@˜d 5§<¡<Ñ0ð ¸5¿<¹<ð ÐRWð Ð\a×\hÑ\hô ð4 ;Bñ	Ø�u—|‘|Ñ$ðà�\‰\ðð ðð "Ð"6Ñ7ð	ð
 ‡\�\ör   