ó
    >:j3-  ã                  óä   • S SK Jr  S SKrS SKrS SKJr  S SKrS SKJr  S SK	J
r
  S SKJr  S SKJr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KJrJr        SS jr " S S\5      rg)é    )ÚannotationsN)ÚUnion)Ú_calculate_correct_fan)ÚConv1D)Úis_bnb_4bit_availableÚis_bnb_available)Ú	BaseTunerÚBaseTunerLayer)Ú2TRANSFORMERS_MODELS_TO_VERA_TARGET_MODULES_MAPPINGé   )Ú
BufferDict)Ú _maybe_include_all_linear_layersé   )Ú
VeraConfig)ÚLinearÚ	VeraLayerc                óˆ  • [        U [        5      (       a  [        R                  " U 5      nOU n[	        US5      n[
        R                  " S5      nU[
        R                  " U5      -  n[
        R                  " S5      U-  n[        R                  " 5          UR                  U* XaS9sSSS5        $ ! , (       d  f       g= f)a˜  
Kaiming Uniform Initialisation adapted to accept a `torch.Generator` object for PRNG.

Args:
    tensor_or_shape (`Union[torch.Tensor, tuple[int, ...]]`):
        Tensor to initialise, or shape of new tensor to create and then initialise.
    generator: (`torch.Generator`):
        Generator object that manages the state of the PRNG algorithm in use.

Returns:
    `torch.Tensor`: The initialised tensor.
Úfan_inr   g      @©Ú	generatorN)	Ú
isinstanceÚtupleÚtorchÚemptyr   ÚmathÚsqrtÚno_gradÚuniform_)Útensor_or_shaper   ÚtensorÚfanÚgainÚstdÚbounds          ÚS/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/tuners/vera/model.pyÚ_kaiming_initr&   &   s†   € ô  �/¤5×)Ñ)Ü—’˜_Ó-‰à ˆÜ
  ¨Ó
2€CÜ�9Š9�Q‹<€DØ
”—’˜3“Ñ
€CÜ�IŠI�c‹N˜SÑ €Eä	�Š�Ø�‰ ˜v uˆÐB÷ 
��ús   ÂB3Â3
Cc                  ó|   ^ • \ rS rSr% SrSrS\S'   \r\	r
SS jrSS jrSS jrSU 4S	 jjrS
 r\S 5       rSrU =r$ )Ú	VeraModeléC   aé  
Creates Vector-based Random Matrix Adaptation (Vera) model from a pretrained transformers model.

Args:
    model ([`~transformers.PreTrainedModel`]): The model to be adapted.
    config ([`VeraConfig`]): The configuration of the Vera model.
    adapter_name (`str`): The name of the adapter, defaults to `"default"`.
    low_cpu_mem_usage (`bool`, `optional`, defaults to `False`):
        Create empty adapter weights on meta device. Useful to speed up the loading process.

Returns:
    `torch.nn.Module`: The Vera model.

Example:

    ```py
    >>> from transformers import AutoModelForCausalLM
    >>> from peft import VeraConfig, get_peft_model

    >>> base_model = AutoModelForCausalLM.from_pretrained("facebook/opt-125m")
    >>> config = VeraConfig(r=128)
    >>> model = get_peft_model(base_model, config)
    ```

**Attributes**:
    - **model** ([`~transformers.PreTrainedModel`]) -- The model to be adapted.
    - **peft_config** ([`VeraConfig`]): The configuration of the Vera model.
Úvera_lambda_ÚstrÚprefixc                ó²  • U R                  U R                  5      nU R                  X5      n[        X0R                  5      nSnU R                  R	                  5        Hå  u  pVU R                  X55      (       d  M  [        U[        R                  5      (       a  UR                  UR                  4nOg[        U[        5      (       aP  [        UR                  S5      (       a  UR                  R                  OUR                  R                  nUSSS2   nOM¼  Uc  UnMÃ  Xt:w  d  MÊ  [!        S [#        XG5       5       5      nMç     Uc  Sn[%        U5      eU$ )z¼
Finds the largest input and output dimensions across linear layers that have been wrapped with VeRA.

This will be used for determining the size of the shared vera_A and vera_B matrices.
NÚds_shapeéÿÿÿÿc              3  ó<   #   • U  H  u  p[        X5      v •  M     g 7f©N)Úmax)Ú.0ÚaÚbs      r%   Ú	<genexpr>Ú&VeraModel._find_dim.<locals>.<genexpr>‚   s   é € Ð%]Ò<\±D°A¤c¨!§i iÒ<\ùs   ‚z[No layers types compatible with VeRA were found. Please check `peft_config.target_modules`.)Úget_model_configÚmodelÚ_prepare_adapter_configr   Únamed_modulesÚ_check_target_module_existsr   Únnr   Úout_featuresÚin_featuresr   ÚhasattrÚweightr.   Úshaper   ÚzipÚ
ValueError)	ÚselfÚconfigÚmodel_configÚpeft_configÚlargest_shapeÚkeyÚmoduleÚmodule_shapeÚmsgs	            r%   Ú	_find_dimÚVeraModel._find_dime   s"  € ð ×,Ñ,¨T¯Z©ZÓ8ˆà×2Ñ2°6ÓHˆÜ6°{ÇJÁJÓOˆàˆØŸ:™:×3Ñ3Ö5‰KˆCØ×3Ñ3°K×EÑEÙä˜&¤"§)¡)×,Ñ,Ø%×2Ñ2°F×4FÑ4FÐF‘Ü˜F¤F×+Ñ+Ü9@ÀÇÁÐPZ×9[Ñ9[˜vŸ}™}×5Ò5Ðag×anÑan×atÑat�Ø+©D¨b¨DÑ1‘áàÑ$Ø ,�ÙàÕ,Ü %Ñ%]¼CÀÔ<\Ó%]Ó ]’ñ# 6ð& Ñ ØoˆCÜ˜S“/Ð!àÐó    c                óv  • U R                  U5      u  p4[        0 UR                  S9U l        [        0 UR                  S9U l        [
        R                  " SS9R                  UR                  5      n[        UR                  U4US9n[        X1R                  4US9nX`R                  U'   XpR                  U'   g )N)Ú
persistentÚcpu)Údevicer   )rN   r   Úsave_projectionÚvera_AÚvera_Br   Ú	GeneratorÚmanual_seedÚprojection_prng_keyr&   Úr)rE   rF   Úadapter_nameÚlinear_out_dimÚlinear_in_dimr   rV   rW   s           r%   Ú_init_vera_A_vera_BÚVeraModel._init_vera_A_vera_BŠ   sž   € Ø(,¯©°vÓ(>Ñ%ˆô ! °×0FÑ0FÑGˆŒÜ  °×0FÑ0FÑGˆŒô —O’O¨5Ñ1×=Ñ=¸f×>XÑ>XÓYˆ	Ü §¡¨-Ð8ÀIÑNˆÜ ·±Ð9ÀYÑOˆà$*�‰�LÑ!Ø$*�‰�LÒ!rP   c                ó&   • U R                  X#5        g r1   )r_   )rE   r9   rF   r\   s       r%   Ú_pre_injection_hookÚVeraModel._pre_injection_hook™   s   € Ø× Ñ  Õ6rP   c                ó²  >• [         TU ]  U5        U R                  R                  5        HJ  nX!L a  M	  UR                  UR                  :w  d  M%  [        SUR                  < SUR                   S35      e   [        U R                  R                  5        Vs1 s H  oR                  iM     sn5      n[        U5      S:”  a  [        SU 35      egs  snf )z´
A helper method to check the config when a new adapter is being added.

Raise a ValueError if there is something wrong with the config or if it conflicts with existing adapters.

z_Vera PRNG initialisation key must be the same for all adapters. Got config.projection_prng_key=z but previous config had Ú.r   zcVeRA projection weights must be saved for all adapters or none, but got multiple different values: N)	ÚsuperÚ_check_new_adapter_configrH   ÚvaluesrZ   rD   ÚsortedrU   Úlen)rE   rF   Úexisting_configÚsave_project_unique_valuesÚ	__class__s       €r%   rg   Ú#VeraModel._check_new_adapter_configœ   sä   ø€ ô 	‰Ñ)¨&Ô1à#×/Ñ/×6Ñ6Ö8ˆOØÒ(áà×2Ñ2°f×6PÑ6PÕPÜ ØvÐ[a×[uÑ[uÑZwð x+Ø+:×+NÑ+NÐ*OÈqðRóð ñ  9ô &,ÐRV×RbÑRb×RiÑRiÔRkÓ,lÒRkÈ×-CÔ-CÑRkÑ,lÓ%mÐ"ÜÐ)Ó*¨QÓ.ÜØuØ-Ð.ð0óð ð /ùò -ms   ÂCc           
     ó”  • Uc  [        S5      eUR                  n[        US5      =(       a    UR                  S Ln	UUR                  UR
                  UR                  [        U R                  SS5      [        U R                  SS5      S.n
XšS'   [        U[        5      (       aH  UR                  UU R                  U R                  UUR                  UR                  UR                  S9  g U R                  " XR                  U R                  X#40 U
D6nX R                   ;  a  UR#                  S5        U R%                  XTX³5        g )NzCurrent Key shouldn't be `None`ÚbiasÚis_loaded_in_8bitFÚis_loaded_in_4bit)r[   Úvera_dropoutÚfan_in_fan_outÚinit_weightsÚloaded_in_8bitÚloaded_in_4bit)Ú	d_initial)rD   r[   r@   rp   rs   rt   ru   Úgetattrr9   r   r   Úupdate_layerrV   rW   rx   Ú_create_new_moduleÚactive_adapterÚrequires_grad_Ú_replace_module)rE   Úvera_configr\   ÚtargetÚtarget_nameÚparentÚcurrent_keyÚoptional_kwargsr[   rp   ÚkwargsÚ
new_modules               r%   Ú_create_and_replaceÚVeraModel._create_and_replace·   s'  € ð ÑÜÐ>Ó?Ð?à�M‰MˆÜ�v˜vÓ&×B¨6¯;©;¸dÐ+BˆàØ'×4Ñ4Ø)×8Ñ8Ø'×4Ñ4Ü% d§j¡jÐ2EÀuÓMÜ% d§j¡jÐ2EÀuÓMñ
ˆð ˆv‰ä�fœf×%Ñ%Ø×ÑØØ—‘Ø—‘ØØ×(Ñ(Ø×(Ñ(Ø%×/Ñ/ð  ò ð ×0Ò0°¿k¹kÈ4Ï;É;ÐXdÑwÐpvÑwˆJØ×#6Ñ#6Ó6à×)Ñ)¨%Ô0Ø× Ñ  °jÕIrP   c                óâ  • [        5       (       a
  SS KnSSKJn  [	        5       (       a  SSKJn  UR                  SS5      n	UR                  SS5      n
UR                  SS5      n[        U[        5      (       a  UR                  5       nOUnU
(       a†  [        UWR                  R                  5      (       aa  UR                  5       nUR                  UR                  R                  UR                  R                   UR"                  S	.5        W" XCX40 UD6$ U(       a†  [        UWR                  R
                  5      (       aa  UR                  5       nUR                  UR$                  UR&                  R(                  UR&                  R*                  S
.5        W" XCX40 UD6$ [        U[,        R                  R.                  5      (       a-  US   (       a"  [0        R2                  " S5        S=US'   U l        OV[        U[6        5      (       a2  SUS'   US   (       d"  [0        R2                  " S5        S=US'   U l        O[9        SU S35      e[/        UUUU4U	U R:                  S.UD6nU$ )Nr   r   )ÚLinear8bitLt)Ú
Linear4bitrp   Frv   rw   )Úhas_fp16_weightsÚ	thresholdÚindex)Úcompute_dtypeÚcompress_statisticsÚ
quant_typert   zjfan_in_fan_out is set to True but the target module is `torch.nn.Linear`. Setting fan_in_fan_out to False.TÚis_target_conv_1d_layerzafan_in_fan_out is set to False but the target module is `Conv1D`. Setting fan_in_fan_out to True.zTarget module z is not supported. Currently, only the following modules are supported: `torch.nn.Linear`, `transformers.pytorch_utils.Conv1D`.)rp   rx   )r   ÚbitsandbytesÚbnbrŠ   r   r‹   ÚpopÚgetr   r
   Úget_base_layerr=   ÚcopyÚupdateÚstaterŒ   r�   rŽ   r�   rA   r�   r‘   r   r   ÚwarningsÚwarnrt   r   rD   rx   )r   rV   rW   r\   r€   r…   r”   rŠ   r‹   rp   rv   rw   Útarget_base_layerÚeightbit_kwargsÚfourbit_kwargsr†   s                   r%   r{   ÚVeraModel._create_new_moduleá   sG  € ô ×ÑÛ&å)ä ×"Ñ"Ý'à�z‰z˜& %Ó(ˆØŸ™Ð$4°eÓ<ˆØŸ™Ð$4°eÓ<ˆä�fœn×-Ñ-Ø &× 5Ñ 5Ó 7Ñà &ÐæœjÐ):¸C¿F¹F×<OÑ<O×PÑPØ$Ÿk™k›mˆOØ×"Ñ"à(9×(?Ñ(?×(PÑ(PØ!2×!8Ñ!8×!BÑ!BØ.×4Ñ4ñôñ   °fÑXÈÑXÐXÞ¤
Ð+<¸c¿f¹f×>OÑ>O× PÑ PØ#Ÿ[™[›]ˆNØ×!Ñ!à%6×%DÑ%DØ+<×+CÑ+C×+WÑ+WØ"3×":Ñ":×"EÑ"Eñôñ ˜f°FÑUÀnÑUÐUÜÐ)¬5¯8©8¯?©?×;Ñ;ØÐ&×'Ü—’ð7ôð INÐM�Ð'Ñ(¨;Ô+EøÜÐ)¬6×2Ñ2Ø04ˆFÐ,Ñ-ØÐ*×+Ü—’Øwôð IMÐL�Ð'Ñ(¨;Ô+EøäØ   ð )Jð Jóð ô ØØØØð	
ð
 Ø!×+Ñ+ñ
ð ñ
ˆ
ð ÐrP   )rV   rW   )Úreturnztuple[int, int])rF   r   r\   r+   r¡   ÚNone)r9   z	nn.ModulerF   r   r\   r+   r¡   r¢   )rF   r   r¡   r¢   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r,   Ú__annotations__r   Útuner_layer_clsr   Útarget_module_mappingrN   r_   rb   rg   r‡   Ústaticmethodr{   Ú__static_attributes__Ú__classcell__)rm   s   @r%   r(   r(   C   sQ   ø‡ ñð: !€FˆCÓ Ø€OØNÐô#ôJ+ô7÷ò6(JðT ñDó öDrP   r(   )r   z$Union[torch.Tensor, tuple[int, ...]]r   ztorch.Generatorr¡   ztorch.Tensor) Ú
__future__r   r   r›   Útypingr   r   Útorch.nnr=   Útorch.nn.initr   Útransformers.pytorch_utilsr   Úpeft.import_utilsr   r   Úpeft.tuners.tuners_utilsr	   r
   Ú
peft.utilsr   Ú_buffer_dictr   Útuners_utilsr   rF   r   Úlayerr   r   r&   r(   © rP   r%   Ú<module>rº      si   ðõ #ã Û Ý ã Ý Ý 0Ý -ç Eß >õõ &Ý ;Ý ß $ðCØ9ðCàðCð ôCô:c�	õ crP   