ó
    >:j¥3  ã                   óÈ   • S SK r S SKJr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  S SKJrJr  S SKJr  SSKJr  SSKJr   " S	 S
\5      r " S S\R,                  \5      rg)é    N)ÚAnyÚOptionalÚUnion)ÚConv1D)ÚBaseTunerLayerÚcheck_adapters_to_merge)Ú	transposeé   )ÚWAVELET_REDUCTIONS)Ú	waverec2dc                   óž   • \ rS rSrSrSrS\R                  SS4S jr SS jr	\
R                  " 5       S	 5       rS\
R                  4S
 jrSrg)ÚWaveFTLayeré   )Úwaveft_spectrum)Úwaveft_n_frequencyÚwaveft_scalingÚwaveft_random_loc_seedÚwaveft_wavelet_familyÚwaveft_indicesÚwaveft_use_idwtÚ
base_layerÚreturnNc                 óh  • Xl         0 U l        0 U l        [        R                  " 0 5      U l        0 U l        0 U l        0 U l        0 U l	        SU l
        / U l        X l        U R                  5       n[        U[        R                  5      (       a$  UR                   UR"                  sU l        U l        g [        U[$        5      (       aU  ['        UR(                  S5      (       a  UR(                  R*                  OUR(                  R,                  u  U l        U l        g [/        S[1        U5       35      e)NFÚds_shapezUnsupported layer type )r   r   r   ÚnnÚParameterDictr   r   r   r   r   Ú_disable_adaptersÚmerged_adaptersÚkwargsÚget_base_layerÚ
isinstanceÚLinearÚin_featuresÚout_featuresr   ÚhasattrÚweightr   ÚshapeÚ
ValueErrorÚtype)Úselfr   r   s      ÚU/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/tuners/waveft/layer.pyÚ__init__ÚWaveFTLayer.__init__+   s  € Ø$ŒØ"$ˆÔØ ˆÔÜ!×/Ò/°Ó3ˆÔØ%'ˆÔ"Ø ˆÔØ&(ˆÔ#Ø!ˆÔà!&ˆÔØ!ˆÔØŒà×(Ñ(Ó*ˆ
Ü�j¤"§)¡)×,Ñ,Ø2<×2HÑ2HÈ*×JaÑJaÐ/ˆDÔ˜dÕ/Ü˜
¤F×+Ñ+ä.5°j×6GÑ6GÈ×.TÑ.T�
×!Ñ!×*Ò*ÐZd×ZkÑZk×ZqÑZqñ 0ˆDÔ˜dÕ/ô Ð6´t¸JÓ7GÐ6HÐIÓJÐJó    c                 óú  • US::  a  [        SU 35      eX R                  U R                  -  :”  a(  [        SU SU R                  U R                  -   35      eX R                  U'   XPR                  U'   X`R
                  U'   XpR                  U'   [        U   u  p‰[        R                  " 5       R                  U R                  U   5      n
[        R                  " U R                  U R                  -  U
S9S U n[        R                  " X°R                  -  X°R                  -  /SS9U R                  U'   X0R                  U'   U(       aH  [        R                   " [        R"                  " U5      SS9U R$                  U'   U R'                  U5        O;S	n[        R                   " [        R(                  " U5      U-  SS9U R$                  U'   U R+                  U5        U R-                  U R.                  5        g )
Nr   zI`n_frequency` should be a positive integer value but the value passed is zu`n_frequency` should be less than or equal to the product of the input and output dimensions but the value passed is z and the product is )Ú	generator)ÚdimT)Úrequires_gradg{®Gáz„?)r(   r#   r$   r   r   r   r   r   ÚtorchÚ	GeneratorÚmanual_seedÚrandpermÚstackr   r   r   Ú	ParameterÚemptyr   Úreset_wave_parametersÚrandnÚ%_move_adapter_to_device_of_base_layerÚset_adapterÚactive_adapters)r*   Úadapter_nameÚn_frequencyÚscalingÚinit_weightsÚrandom_loc_seedÚwavelet_familyÚuse_idwtÚreduction_rowsÚreduction_colsr0   ÚindicesÚstd_devs                r+   Úupdate_layerÚWaveFTLayer.update_layerC   sÚ  € ð ˜!ÓÜÐhÐitÐhuÐvÓwÐwØ×)Ñ)¨D×,=Ñ,=Ñ=Ó=Üð+Ø+6¨-Ð7KÈD×L\ÑL\Ð_c×_pÑ_pÑLpÐKqðsóð ð
 1<×Ñ Ñ-Ø4C×#Ñ# LÑ1Ø3A×"Ñ" <Ñ0Ø-5×Ñ˜\Ñ*ô *<¸NÑ)KÑ&ˆô —O’OÓ%×1Ñ1°$×2MÑ2MÈlÑ2[Ó\ˆ	Ü—.’. ×!2Ñ!2°T×5EÑ5EÑ!EÐQZÑ[Ð\hÐ]hÐiˆô -2¯KªKØ×(Ñ(Ñ(¨'×4DÑ4DÑ*DÐEÈ1ñ-
ˆ×Ñ˜LÑ)ð -4×Ñ˜LÑ)ö ä13·²¼e¿kºkÈ+Ó>VÐfjÑ1kˆD× Ñ  Ñ.Ø×&Ñ& |Õ4ð ˆGÜ13·²¼e¿kºkÈ+Ó>VÐY`Ñ>`ÐptÑ1uˆD× Ñ  Ñ.à×2Ñ2°<Ô@Ø×Ñ˜×-Ñ-Õ.r.   c                 ó˜   • XR                   R                  5       ;   a-  [        R                  R	                  U R                   U   5        g g ©N)r   Úkeysr   ÚinitÚzeros_©r*   r?   s     r+   r:   Ú!WaveFTLayer.reset_wave_parametersp   s7   € à×/Ñ/×4Ñ4Ó6Ó6Ü�G‰G�N‰N˜4×/Ñ/°Ñ=Õ>ð 7r.   c                 óŒ  • U R                   U   nU R                  U   R                  UR                  5      nU R                  U   nU R
                  U   (       Ga  [        U   u  pVU R                  U-   nU R                  U-   nUS-  S:w  a  US-  nUS-  S:w  a  US-  n[        R                  " XxUR                  UR                  S9n	XpR                  -
  S-  n
X€R                  -
  S-  nUR                  5       nUSS S 24==   U
-  ss'   USS S 24==   U-  ss'   USS S 24   U:  USS S 24   U:  -  nUS S 2U4   nX-   nXùUSS S 24   USS S 24   4'   U	R                  u  nnUS-  US-  nnU	S U2S U24   nU	S U2US 24   nU	US 2S U24   nU	US 2US 24   nUUUU44n[        UU5      U R                  U   -  nUR                  S   U R                  :w  d  UR                  S   U R                  :w  ac  UR                  S   U R                  -
  S-  nUR                  S   U R                  -
  S-  nUUUU R                  -   2UUU R                  -   24   nU$ [        R                  " U R                  U R                  UR                  UR                  S9n	X)USS S 24   USS S 24   4'   X�R                  U   -  nU$ )Né   r   r
   )ÚdeviceÚdtype)r   r   ÚtorU   r   r   r   r$   r#   r3   ÚzerosrV   Úcloner'   r   r   )r*   ÚadapterÚspectrumrH   rD   rF   rG   Úpadded_out_featuresÚpadded_in_featuresÚdense_spectrumÚ
row_offsetÚ
col_offsetÚpadded_indicesÚ
valid_maskÚvalid_indicesÚvalid_spectrumÚHÚWÚH2ÚW2ÚcAÚcHÚcVÚcDÚcoeffsÚdelta_weightÚ	start_rowÚ	start_cols                               r+   Úget_delta_weightÚWaveFTLayer.get_delta_weightu   sZ  € Ø×'Ñ'¨Ñ0ˆØ×%Ñ% gÑ.×1Ñ1°(·/±/ÓBˆØ×3Ñ3°GÑ<ˆð ×Ñ ×(Ð(Ü-?ÀÑ-OÑ*ˆNð #'×"3Ñ"3°nÑ"DÐØ!%×!1Ñ!1°NÑ!BÐð # QÑ&¨!Ó+Ø# qÑ(Ð#Ø! AÑ%¨Ó*Ø" aÑ'Ð"ô #Ÿ[š[Ø#ÀÇÁÐW_×WeÑWeñˆNð
 .×0AÑ0AÑAÀaÑGˆJØ,×/?Ñ/?Ñ?ÀAÑEˆJð %Ÿ]™]›_ˆNØ˜1ša˜4Ó  JÑ.Ó Ø˜1ša˜4Ó  JÑ.Ó ð )¨ªA¨Ñ.Ð1DÑDÈÐXYÒ[\ÐX\ÑI]Ð`rÑIrÑsˆJØ*ª1¨j¨=Ñ9ˆMØ%Ñ1ˆNð HV˜=¨ªA¨Ñ.°¸aÂ¸dÑ0CÐCÑDð "×'Ñ'‰DˆAˆqØ˜!‘V˜Q !™V�ˆBØ    S b S Ñ)ˆBØ    R¡S Ñ)ˆBØ ¡ S b S Ñ)ˆBØ ¡ R¡S Ñ)ˆBð ˜2˜r 2˜,Ð'ˆFô % V¨^Ó<¸t×?RÑ?RÐSZÑ?[Ñ[ˆLð ×!Ñ! !Ñ$¨×(9Ñ(9Ó9¸\×=OÑ=OÐPQÑ=RÐVZ×VfÑVfÓ=fà)×/Ñ/°Ñ2°T×5FÑ5FÑFÈ1ÑL�	Ø)×/Ñ/°Ñ2°T×5EÑ5EÑEÈ!ÑK�	ð  ,Ø 	¨D×,=Ñ,=Ñ =Ð=¸yÈ9ÐW[×WgÑWgÑKgÐ?gÐgñ �ð Ðô #Ÿ[š[Ø×!Ñ! 4×#3Ñ#3¸H¿O¹OÐS[×SaÑSañˆNð <D˜7 1¢a 4™=¨'°!²Q°$©-Ð7Ñ8Ø)×,?Ñ,?ÀÑ,HÑHˆLàÐr.   )r   r   r#   r   r   r$   r   r   r   r   r   r   r   )Údb1T)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Úadapter_layer_namesÚother_param_namesr   ÚModuler,   rJ   r3   Úno_gradr:   ÚTensorrq   Ú__static_attributes__© r.   r+   r   r      sc   † à.ÐðÐðK 2§9¡9ð K¸4ô Kð2 quô+/ðZ ‡]‚]ƒ_ñ?ó ð?ðK¨5¯<©<÷ Kr.   r   c                   ó  ^ • \ rS rSr       SS\S\S\S\S\\\4   S\S\S	\S
S4U 4S jjjr	SS\S\
\\      S
S4S jjrSS jrS\R                  S\S\S
\R                  4S jrSS\S
\4S jjrS
\4U 4S jjrSrU =r$ )ÚWaveFTLinearéÃ   r?   r@   rA   Úfan_in_fan_outrB   rC   rD   rE   r   Nc
           	      ó�   >• [         TU ]  5         [        R                  " X40 U
D6  XPl        X l        U R                  X#XFXxU	5        g rM   )Úsuperr,   r   r‚   Ú_active_adapterrJ   )r*   r   r?   r@   rA   r‚   rB   rC   rD   rE   r   Ú	__class__s              €r+   r,   ÚWaveFTLinear.__init__Å   sD   ø€ ô 	‰ÑÔÜ×Ò˜TÑ8°Ò8Ø,ÔØ+ÔØ×Ñ˜,°WÈOÐmuÕvr.   Ú
safe_mergeÚadapter_namesc                 ó�  • [        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[        U R                  U5      U R                  5      -  n[        R                  " U5      R                  5       (       d  [        SU S35      eXTR                  l        OBUR                  =R
                  [        U R                  U5      U R                  5      -  s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   r   rN   r    r&   ÚdatarY   r	   rq   r‚   r3   ÚisfiniteÚallr(   r   Úappend)r*   rˆ   r‰   Úactive_adapterr   Úorig_weightss         r+   ÚmergeÚWaveFTLinear.mergeØ   s	  € ô 0°ÓDˆÞàä+ˆNØ×!5Ñ!5×!:Ñ!:Ó!<Õ<Ø!×0Ñ0Ó2�
Þð $.×#4Ñ#4×#9Ñ#9×#?Ñ#?Ó#A�LØ ¤I¨d×.CÑ.CÀNÓ.SÐUY×UhÑUhÓ$iÑi�Lä Ÿ>š>¨,Ó7×;Ñ;×=Ñ=Ü(ØOÐP^ÐO_Ð_rÐsóð ð .:×%Ñ%Õ*à×%Ñ%×*Ò*¬i¸×8MÑ8MÈnÓ8]Ð_c×_rÑ_rÓ.sÑsÕ*Ø×$Ñ$×+Ñ+¨N×;ò# ,r.   c                 óÌ  • U R                   (       d  [        R                  " S5        g[        U R                  5      S:”  a£  U R                  R                  5       nXR                  R                  5       ;   aP  U R                  5       R                  =R                  [        U R                  U5      U R                  5      -  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   )ÚmergedÚwarningsÚwarnÚlenr   Úpopr   rN   r    r&   r‹   r	   rq   r‚   )r*   r�   s     r+   ÚunmergeÚWaveFTLinear.unmergeý   s©   € ð �{�{Ü�MŠMÐ<Ô=ØÜ�$×&Ñ&Ó'¨!Ó+Ø!×1Ñ1×5Ñ5Ó7ˆNØ×!5Ñ!5×!:Ñ!:Ó!<Ó<Ø×#Ñ#Ó%×,Ñ,×1Ò1´YØ×)Ñ)¨.Ó9¸4×;NÑ;Nó6ñ Õ1ô �$×&Ñ&Ó'¨!×+r.   ÚxÚargsr   c                 ó<  • UR                   nU R                  (       a8  U R                  (       a  U R                  5         U R                  " U/UQ70 UD6nOµU R                  (       a  U R                  " U/UQ70 UD6nO�U R                  " U/UQ70 UD6nU R
                   Hg  nX`R                  R                  5       ;  a  M"  U R                  U5      nU R                  XR                   5      nU[        R                  " X5      -   nMi     UR                  U5      nU$ rM   )rV   Údisable_adaptersr”   r™   r   r>   r   rN   rq   Ú_cast_input_dtypeÚFÚlinearrW   )r*   r›   rœ   r   Úprevious_dtypeÚresultr�   Údelta_ws           r+   ÚforwardÚWaveFTLinear.forward  sç   € ØŸ™ˆà× × Ø�{�{Ø—‘”Ø—_’_ QÐ8¨Ò8°Ñ8‰FØ�[�[Ø—_’_ QÐ8¨Ò8°Ñ8‰Fà—_’_ QÐ8¨Ò8°Ñ8ˆFØ"&×"6Ô"6�Ø!×)=Ñ)=×)BÑ)BÓ)DÓDÙà×/Ñ/°Ó?�Ø×*Ñ*¨1¯m©mÓ<�Ø¤!§(¢(¨1Ó"6Ñ6’ñ #7ð —‘˜>Ó*ˆØˆr.   c                 óL   • [        U R                  5       [        5      (       a  gg)NFT)r!   r    r   rQ   s     r+   Úsupports_lora_conversionÚ%WaveFTLinear.supports_lora_conversion!  s    € Ü�d×)Ñ)Ó+¬V×4Ñ4ð Ør.   c                 ó*   >• [         TU ]  5       nSU-   $ )Nzwaveft.)r„   Ú__repr__)r*   Úrepr†   s     €r+   r«   ÚWaveFTLinear.__repr__(  s   ø€ Ü‰gÑÓ ˆØ˜3‰Ðr.   )r…   r‚   )iè  g     Àb@FFi	  rs   T)FN)r   N)Údefault)rt   ru   rv   rw   ÚstrÚintÚfloatÚboolr   r,   r   Úlistr‘   r™   r3   r|   r   r¥   r¨   r«   r}   Ú__classcell__)r†   s   @r+   r€   r€   Ã   s	  ø† ð  ØØ$Ø).Ø"Ø#Øñwð ðwð ð	wð
 ðwð ðwð ˜D #˜IÑ&ðwð ðwð ðwð ðwð 
÷wð wñ&#< ð #<¸XÀdÈ3ÁiÑ=Pð #<Ð\`õ #<ôJð˜Ÿ™ð ¨cð ¸Sð ÀUÇ\Á\ô ñ,°Sð Èõ ð˜#÷ õ r.   r€   )r•   Útypingr   r   r   r3   Útorch.nnr   Útorch.nn.functionalÚ
functionalr    Útransformers.pytorch_utilsr   Úpeft.tuners.tuners_utilsr   r   Úpeft.utils.otherr	   Ú	constantsr   r   r   rz   r€   r~   r.   r+   Ú<module>r½      sN   ðó ß 'Ñ 'ã Ý ß Ð Ý -ç LÝ &å )Ý  ôb�.ô bôJg�2—9‘9˜kõ gr.   