ó
    >:j<  ã                  ób   • S SK Jr  S SKrS SKrS SKJrJr  S SKJr  SSK	J
r
Jr   " S S\5      rg)	é    )ÚannotationsN)Ú	BaseTunerÚBaseTunerLayer)Ú3TRANSFORMERS_MODELS_TO_SHIRA_TARGET_MODULES_MAPPINGé   )ÚLinearÚ
ShiraLayerc                  óF   • \ rS rSr% SrSrS\S'   \r\	r
S r\S 5       rSrg	)
Ú
ShiraModelé   a:  
Creates a Sparse High Rank Adapter (SHiRA) Model from a pretrained model.

Args:
    model ([`~transformers.PreTrainedModel`]): The model to be adapted.
    config ([`ShiraConfig`]): The configuration of the SHiRA model.
    adapter_name (`str`): The name of the adapter, defaults to `"default"`.

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

Example:

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

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

**Attributes**:
    - **model** ([`~transformers.PreTrainedModel`]) -- The model to be adapted.
    - **peft_config** ([`ShiraConfig`]): The configuration of the SHiRA model.
Úshira_ÚstrÚprefixc                óX  • Uc  [        S5      e[        US5      =(       a    UR                  S Ln0 n	X‰S'   UR                  S:X  a  UR                  U	S'   UR                  5        H	  u  p«X¹U
'   M     [        U[        5      (       a^  UR                  b(  UR                  " UR                  UR                  40 U	D6OS nUR                  UUUR                  UR                  S9  g U R                  " XU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ÚrandomÚrandom_seed)Úinit_weightsF)Ú
ValueErrorÚhasattrr   Ú	mask_typer   ÚitemsÚ
isinstancer   Úmask_fnÚ
base_layerÚrÚupdate_layerr   Ú_create_new_moduleÚactive_adapterÚrequires_grad_Ú_replace_module)ÚselfÚshira_configÚadapter_nameÚtargetÚtarget_nameÚparentÚcurrent_keyÚoptional_kwargsr   ÚkwargsÚkÚvÚmaskÚ
new_modules                 ÚT/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/tuners/shira/model.pyÚ_create_and_replaceÚShiraModel._create_and_replace=   s+  € ð ÑÜÐ>Ó?Ð?ä�v˜vÓ&×B¨6¯;©;¸dÐ+BˆØˆØˆv‰Ø×!Ñ! XÓ-Ø$0×$<Ñ$<ˆF�=Ñ!à#×)Ñ)Ö+‰DˆAØ�1‹Iñ ,ô �fœf×%Ñ%ð  ×'Ñ'Ñ3ð ×$Ò$ V×%6Ñ%6¸¿¹ÑQÈ&ÒQàð ð
 ×ÑØØØ—‘Ø)×6Ñ6ð	  ò ð ×0Ò0°ÈVÑ^ÐW]Ñ^ˆJØ×#6Ñ#6Ó6à×)Ñ)¨%Ô0Ø× Ñ  °jÕIó    c                óò  • U R                   nUR                  SS5      n[        U[        5      (       a  UR	                  5       nOUn[        U[
        R                  R                  5      (       a&  U(       a  [        R                  " S5        S=o@l         O[        SU S35      eU R                  b  U R                  " X`R                  40 UD6OS n[        UUUU R                  U4SU R                  0UD6nU$ )Nr   Fzjfan_in_fan_out is set to True but the target module is `torch.nn.Linear`. Setting fan_in_fan_out to False.zTarget module zZ is not supported. Currently, only the following modules are supported: `torch.nn.Linear`.r   )Úfan_in_fan_outÚpopr   r   Úget_base_layerÚtorchÚnnr   ÚwarningsÚwarnr   r   r   r   )	r#   r$   r%   r*   r4   Ú_Útarget_base_layerr-   r.   s	            r/   r   ÚShiraModel._create_new_modulef   s  € à%×4Ñ4ˆà�J‰J�v˜uÓ%ˆä�fœn×-Ñ-Ø &× 5Ñ 5Ó 7Ñà &ÐäÐ'¬¯©¯©×9Ñ9ÞÜ—’ð7ôð @EÐD�Ô!<øäØ   ð )%ð %óð ð ×#Ñ#Ñ/ð × Ò Ð!2·N±NÑMÀfÒMàð 	ô ØØØØ�N‰NØñ
ð &×2Ñ2ð
ð ñ
ˆ
ð Ðr2   © N)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Ú__annotations__r	   Útuner_layer_clsr   Útarget_module_mappingr0   Ústaticmethodr   Ú__static_attributes__r>   r2   r/   r   r      s9   ‡ ñð6 €FˆCÓØ €OØOÐò'JðR ñ'ó ó'r2   r   )Ú
__future__r   r9   r7   Úpeft.tuners.tuners_utilsr   r   Ú
peft.utilsr   Úlayerr   r	   r   r>   r2   r/   Ú<module>rM      s+   ðõ #ã ã ç >õ÷ &ôq�õ qr2   