ó
    >:j!  ã                  ó†   • S SK Jr  S SKrS SKrS SK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Jr   " S	 S
\	5      rg)é    )ÚannotationsN)ÚConv1D)Ú	BaseTunerÚBaseTunerLayer)Ú4TRANSFORMERS_MODELS_TO_VBLORA_TARGET_MODULES_MAPPINGé   )ÚVBLoRAConfig)ÚLinearÚVBLoRALayerc                  ór   • \ rS rSr% SrSrS\S'   \r\	r
SS jrSS jrS r\S	 5       rSSS
 jjrSS jrSrg)ÚVBLoRAModelé   a  
Creates VBLoRA model from a pretrained transformers model.

The method is described in detail in https://huggingface.co/papers/2405.15179.

Args:
    model ([`~transformers.PreTrainedModel`]): The model to be adapted.
    config ([`VBLoRAConfig`]): The configuration of the VBLoRA 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 VBLoRA model.

Example:

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

    >>> base_model = AutoModelForCausalLM.from_pretrained("facebook/opt-125m")
    >>> config = VBLoRAConfig(
    ...     task_type="SEQ_CLS",
    ...     r=4,
    ...     target_modules=["fc1", "fc2", "k_proj", "out_proj", "q_proj", "v_proj"],
    ...     num_vectors=60,
    ...     vector_length=256,
    ...     save_only_topk_weights=True,
    ... )
    >>> model = get_peft_model(base_model, config)
    ```

**Attributes**:
    - **model** ([`~transformers.PreTrainedModel`]) -- The model to be adapted.
    - **peft_config** ([`VBLoRAConfig`]): The configuration of the VBLoRAConfig model.
Úvblora_ÚstrÚprefixc                óô   • [         R                  " UR                  UR                  5      n[         R                  R
                  R                  X1R                  * UR                  5        X0R                  U'   g ©N)	ÚtorchÚzerosÚnum_vectorsÚvector_lengthÚnnÚinitÚuniform_Úinit_vector_bank_boundÚvblora_vector_bank)ÚselfÚconfigÚadapter_namer   s       ÚU/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/tuners/vblora/model.pyÚ_init_vblora_vector_bankÚ$VBLoRAModel._init_vblora_vector_bankH   sV   € Ü"Ÿ[š[¨×);Ñ);¸V×=QÑ=QÓRÐÜ�‰�‰×ÑÐ1×4QÑ4QÐ3QÐSY×SpÑSpÔqØ0B×Ñ Ò-ó    c                ó:   • [         R                  " 0 5      U l        g r   )r   ÚParameterDictr   )r   Úmodelr   r   s       r    Ú_pre_injection_hookÚVBLoRAModel._pre_injection_hookM   s   € Ü"$×"2Ò"2°2Ó"6ˆÕr#   c                ó,  • Uc  [        S5      e[        US5      =(       a    UR                  S LnUR                  US.nU R	                  X5        [        U[        5      (       a]  UR                  UU R                  UR                  UR                  UR                  UR                  UR                  UR                  S9  g U R                  " SUU R                  UUS.UD6n	X R                   ;  a  U	R#                  S5        U R%                  XTX“5        g )NzCurrent Key shouldn't be `None`Úbias)Úfan_in_fan_outr*   )r   r   ÚrÚtopkr   r   Úvblora_dropoutÚinit_logits_std)Úvblora_configr   r   ÚtargetF© )Ú
ValueErrorÚhasattrr*   r+   r!   Ú
isinstancer
   Úupdate_layerr   r,   r-   r   r   r.   r/   Ú_create_new_moduleÚactive_adapterÚrequires_grad_Ú_replace_module)
r   r0   r   r1   Útarget_nameÚparentÚcurrent_keyr*   ÚkwargsÚ
new_modules
             r    Ú_create_and_replaceÚVBLoRAModel._create_and_replaceP   s  € ð ÑÜÐ>Ó?Ð?ä�v˜vÓ&×B¨6¯;©;¸dÐ+Bˆà+×:Ñ:Øñ
ˆð 	×%Ñ% mÔBô �fœf×%Ñ%Ø×ÑØ)Ø#'×#:Ñ#:Ø—/‘/Ø"×'Ñ'Ø)×5Ñ5Ø+×9Ñ9Ø,×;Ñ;Ø -× =Ñ =ð  ò 	ð ×0Ò0ð Ø+Ø#'×#:Ñ#:Ø)Øñ	ð
 ñˆJð ×#6Ñ#6Ó6à×)Ñ)¨%Ô0Ø× Ñ  °jÕIr#   c                óP  • [        U[        5      (       a  UR                  5       nOUn[        U[        R                  R
                  5      (       a-  US   (       a"  [        R                  " S5        S=US'   U l        OV[        U[        5      (       a2  SUS'   US   (       d"  [        R                  " S5        S=US'   U l        O[        SU S35      e[        S
UUUU R                  U R                  U R                  U R                  U R                  U R                   S	.	UD6nU$ )Nr+   zjfan_in_fan_out is set to True but the target module is `torch.nn.Linear`. Setting fan_in_fan_out to False.FTÚ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`.)	Ú
base_layerr   r   r,   r   r   r-   r.   r/   r2   )r5   r   Úget_base_layerr   r   r
   ÚwarningsÚwarnr+   r   r3   r,   r   r   r-   r.   r/   )r0   r   r   r1   r>   Útarget_base_layerr?   s          r    r7   ÚVBLoRAModel._create_new_module|   s5  € ä�fœn×-Ñ-Ø &× 5Ñ 5Ó 7Ñà &ÐäÐ'¬¯©¯©×9Ñ9ØÐ&×'Ü—’ð7ôð KPÐO�Ð'Ñ(¨=Ô+GøÜÐ)¬6×2Ñ2Ø04ˆFÐ,Ñ-ØÐ*×+Ü—’Øwôð KOÐN�Ð'Ñ(¨=Ô+GøäØ   ð )Jð Jóð ô ð 
ØØ1Ø%Ø�o‰oØ%×1Ñ1Ø'×5Ñ5Ø×#Ñ#Ø(×7Ñ7Ø)×9Ñ9ñ
ð ñ
ˆ
ð Ðr#   c                ó²  • SnSnSnU R                  5        H^  u  pVSU;   a  X&R                  5       -  nM  SU;   a  X6R                  5       -  nM9  UR                  (       d  ML  XFR                  5       -  nM`     U R                  U   R                  (       a»  U R                  U   R
                  nSnUS:  a  SnOUS:  a  SnOUS	:  a  SnOS
nX R                  U   R
                  -  U R                  U   R                  S-
  -  n	X R                  U   R
                  -  U R                  U   R                  -  U-  n
[        X9-   U
-   5      nX´4$ X2-   nX´4$ )zP
Returns the number of savable VB-LoRA parameters and other savable parameters.
r   Úvblora_logitsr   r   é   g      Ð?i €  g      à?l        é   )Únamed_parametersÚnumelÚrequires_gradÚpeft_configÚsave_only_topk_weightsr   r-   Úint)r   ÚadapterÚlogits_paramsÚvector_bank_paramsÚother_paramsÚnameÚparamr   ÚfactorÚtopk_weight_paramsÚtopk_indices_paramsÚvblora_paramss               r    Úget_nb_savable_parametersÚ%VBLoRAModel.get_nb_savable_parameters¥   su  € ð ˆØÐØˆØ×0Ñ0Ö2‰KˆDØ $Ó&Ø§¡£Ñ.’Ø%¨Ó-Ø"§k¡k£mÑ3Ò"Ø×$×$Ñ$Ø§¡£Ñ-’ñ 3ð ×Ñ˜GÑ$×;×;Ø×*Ñ*¨7Ñ3×?Ñ?ˆKØˆFØ˜TÓ!Ø‘Ø˜uÓ$Ø‘Ø˜uÓ$Ø‘à�à× 0Ñ 0°Ñ 9× EÑ EÑEÈ×IYÑIYÐZaÑIb×IgÑIgÐjkÑIkÑlð ð × 0Ñ 0°Ñ 9× EÑ EÑEÈ×HXÑHXÐY`ÑHa×HfÑHfÑfÐioÑoð  ô  Ð 2Ñ GÐJ]Ñ ]Ó^ˆMð Ð*Ð*ð /Ñ>ˆMØÐ*Ð*r#   c                óR   • U R                  5       u  p[        SUS SX-   S 35        g)zO
Prints the number of savable VB-LoRA parameters and total savable parameters.
z1VB-LoRA params to-be-saved (float32-equivalent): z,dz || total params to-be-saved: N)r^   Úprint)r   r]   rW   s      r    Úprint_savable_parametersÚ$VBLoRAModel.print_savable_parametersÉ   s>   € ð '+×&DÑ&DÓ&FÑ#ˆÜØ?ÀÈbÐ?Qð R,Ø-:Ñ-IÈ2Ð+NðPõ	
r#   )r   N)r   r	   r   r   ÚreturnÚNone)r&   z	nn.Moduler   r	   r   r   rd   re   )Údefault)rd   ztuple[int, int])rd   re   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Ú__annotations__r   Útuner_layer_clsr   Útarget_module_mappingr!   r'   r@   Ústaticmethodr7   r^   rb   Ú__static_attributes__r2   r#   r    r   r      sQ   ‡ ñ$ðL €FˆCÓØ!€OØPÐôCô
7ò*JðX ñ&ó ð&öP"+÷H
r#   r   )Ú
__future__r   rF   r   Útorch.nnr   Útransformers.pytorch_utilsr   Úpeft.tuners.tuners_utilsr   r   Ú
peft.utilsr   r   r	   Úlayerr
   r   r   r2   r#   r    Ú<module>rw      s0   ðõ #ã ã Ý Ý -ç >Ý Kå  ß &ôt
�)õ t
r#   