ó
    >:jÇ.  ã                  ó‚   • S SK Jr  S SKJr  S SKrS SKJr  S SKJr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)é    )Úannotations)ÚOptionalN)Únn)Ú	BaseTunerÚBaseTunerLayerÚcheck_target_module_exists)Ú4TRANSFORMERS_MODELS_TO_ADAMSS_TARGET_MODULES_MAPPINGé   )ÚAdamssConfig)ÚAdamssLayerÚLinearc                  óØ   ^ • \ rS rSr% SrSrS\S'   \4r\	r
 S     SU 4S jjjr\S 5       r            SS jrSS	 jr        SS
 jrSS jrSS jrSS jrSS jrSrU =r$ )ÚAdamssModelé   a@  
Creates Adamss (Adaptive Multi-Subspaces) model from a pretrained model.

The method decomposes weight matrices using SVD and clusters the decomposed space into multiple trainable subspaces
for parameter-efficient fine-tuning.

Args:
    model (`torch.nn.Module`): The model to be adapted.
    config (`AdamssConfig`): The configuration of the Adamss model.
    adapter_name (`str`): The name of the adapter, defaults to `"default"`.

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

Example:
    ```python
    >>> from transformers import AutoModelForImageClassification
    >>> from peft import AdamssConfig, get_peft_model

    >>> config = AdamssConfig(
    ...     r=500,
    ...     num_subspaces=5,
    ...     target_modules=["query", "value"],
    ... )

    >>> model = AutoModelForImageClassification.from_pretrained("google/vit-base-patch16-224")
    >>> adamss_model = get_peft_model(model, config)
    ```

**Attributes**:
    - **model** ([`~torch.nn.Module`]) -- The model to be adapted.
    - **peft_config** ([`AdamssConfig`]): The configuration of the Adamss model.
Úadamss_ÚstrÚprefixc                ó2   >• 0 U l         [        TU ]	  XX4US9  g )N)Úlow_cpu_mem_usageÚ
state_dict)Ú_asa_total_subspacesÚsuperÚ__init__)ÚselfÚmodelÚconfigÚadapter_namer   r   Ú	__class__s         €ÚU/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/tuners/adamss/model.pyr   ÚAdamssModel.__init__F   s"   ø€ ð %'ˆÔ!Ü‰Ñ˜¨ÐfpÐÒqó    c                ó   • [        X5      $ )z5Helper to check if target module matches the pattern.)r   )Úadamss_configÚkeys     r   Ú_check_target_module_existsÚ'AdamssModel._check_target_module_existsM   s   € ô *¨-Ó=Ð=r!   c           
     óÞ  • Uc  [        S5      e[        U[        5      (       a†  UR                  UUR                  UR
                  UR                  UR                  UR                  UR                  S9  X R                  ;  a%  UR                  U R                  UR                  S9  ggU R                  " XU4SUR                  0UD6nU R                  X!U5        U R                  XTXƒ5        g)z>
Create and replace target module with Adamss-adapted module.
NzCurrent Key shouldn't be `None`)Úinference_moder(   )Ú
ValueErrorÚ
isinstancer   Úupdate_layerÚrÚnum_subspacesÚsubspace_rankÚinit_weightsÚuse_asar(   Úactive_adaptersÚset_adapterÚ_create_new_moduleÚ_record_asa_total_subspacesÚ_replace_module)	r   r#   r   ÚtargetÚtarget_nameÚparentÚcurrent_keyÚoptional_kwargsÚ
new_modules	            r   Ú_create_and_replaceÚAdamssModel._create_and_replaceR   sð   € ð ÑÜÐ>Ó?Ð?ô �fœk×*Ñ*Ø×ÑØØ—‘Ø×+Ñ+Ø×+Ñ+Ø×*Ñ*Ø×%Ñ%Ø,×;Ñ;ð  ñ ð ×#7Ñ#7Ó7à×"Ñ" 4×#7Ñ#7È×HdÑHdÐ"Òeð 8ð
 ×0Ò0Ø¨VñØDQ×D`ÑD`ðØdsñˆJð ×,Ñ,¨\È*ÔUà× Ñ  °jÕIr!   c                óÎ   • [        US5      (       d  gUR                  R                  US5      nUb  US::  a  gU R                  R                  US5      nXT-   U R                  U'   g)zITrack total subspaces per adapter so the ASA schedule matches adamss_pkg.r-   Nr   )Úhasattrr-   Úgetr   )r   r   r#   ÚmoduleÚlayer_num_subspacesÚ
prev_totals         r   r4   Ú'AdamssModel._record_asa_total_subspaces{   sg   € ä�v˜×/Ñ/ØØ$×2Ñ2×6Ñ6°|ÀQÓGÐØÑ&Ð*=ÀÓ*BØØ×.Ñ.×2Ñ2°<ÀÓCˆ
Ø2<Ñ2Rˆ×!Ñ! ,Ò/r!   c                ó|  • [        U[        5      (       a  UR                  5       nOUn[        U[        R                  R
                  5      (       a]  [        UU4UR                  UR                  UR                  UR                  UR                  UR                  UR                  S.UD6nU$ [        SU S35      e)z=
Create a new Adamss module based on the target module type.
)r,   r-   r.   r/   r0   Úuse_dynamic_rankÚsvd_thresholdzTarget module zB is not supported. Currently, only `torch.nn.Linear` is supported.)r*   r   Úget_base_layerÚtorchr   r   r,   r-   r.   r/   r0   rF   rG   Ú	TypeError)r   r#   r   r6   ÚkwargsÚtarget_base_layerr;   s          r   r3   ÚAdamssModel._create_new_module…   s¾   € ô �fœn×-Ñ-Ø &× 5Ñ 5Ó 7Ñà &ÐäÐ'¬¯©¯©×9Ñ9ÜØØðð  —/‘/Ø+×9Ñ9Ø+×9Ñ9Ø*×7Ñ7Ø%×-Ñ-Ø!.×!?Ñ!?Ø+×9Ñ9ñð ñˆJð" Ðô	 Ø   Ð(jÐkóð r!   c                ó¸  • U R                    GHD  nU R                  U   nUR                  (       d  M&  UR                  Us=:*  =(       a    UR                  :*  Os  nU R
                  R                  5        Vs/ s H-  n[        U[        5      (       d  M  X%R                  ;   d  M+  UPM/     nnU(       d  Mª  U(       a/  U H)  nUR                  X#R                  UR                  5        M+     XR                  -  S:H  nU(       d  Mú  U(       d  GM  U R                  X5      n	U	b  U R                  X’U5        U H  nUR!                  U5        M     GMG     gs  snf )uK  
Update importance scores and apply ASA masking (if enabled).

This method should be called in **every** training step after ``loss.backward()`` and before
``optimizer.zero_grad()`` when ASA is enabled. Internally it:

1. Accumulates importance scores via EMA every step during the warmup period.
2. At mask intervals, applies global top-K masking and resets importance.

This is the single entry point for ASA â€“ using the :class:`AdamssAsaCallback` with HuggingFace ``Trainer``
simply delegates to this method. For custom training loops, call this directly instead of the callback.

Args:
    global_step (`int`): The current training step.

Example::

    for step, batch in enumerate(dataloader):
        loss = model(**batch).loss loss.backward() optimizer.step() model.base_model.update_and_allocate(step)
        optimizer.zero_grad()
r   N)r1   Úpeft_configr0   Úinit_warmupÚfinal_warmupr   Úmodulesr*   r   Úexp_avg_ipt_AÚupdate_importanceÚasa_importance_betaÚasa_uncertainty_betaÚmask_intervalÚ_schedule_thresholdÚ_global_mask_to_targetÚreset_importance)
r   Úglobal_stepr   r   Úwithin_warmupÚmÚ
asa_layersrA   Úis_mask_intervalÚcurrent_targets
             r   Úupdate_and_allocateÚAdamssModel.update_and_allocate¨   s)  € ð, !×0Õ0ˆLØ×%Ñ% lÑ3ˆFØ—>—>Ùà"×.Ñ.°+×TÓTÀ×ATÑATÔTˆMð  Ÿ:™:×-Ñ-Ô/óÚ/�a´:¸aÄ×3M“ÐR^×bqÑbqÑRq—Ñ/ð ð ö Ùö Û(�FØ×,Ñ,¨\×;UÑ;UÐW]×WrÑWrÖsñ )ð  +×-AÑ-AÑAÀQÑFÐßˆ}×!1Ñ!1Ø!%×!9Ñ!9¸+Ó!N�Ø!Ñ-Ø×/Ñ/°ÈjÔYó )�FØ×+Ñ+¨LÖ9ô )ò7 1ùòs   Á8EÂEÂ&Ec                óò  • U R                   R                  [        [        U R                  5      S5      S5      nUS:X  a  U R                  5       nUS:X  a  gXR                  :  a  gXR                  ::  aw  SXR                  -
  UR                  UR                  -
  -  -
  n[        S[        SU5      5      n[        USS5      n[        UR                  X2R                  -
  XE-  -  -   5      $ UR                  $ )z<Calculate current target subspaces based on warmup schedule.Údefaultr   Ng      ð?g        Úasa_schedule_exponentg      @)r   r@   ÚnextÚiterr1   Ú_get_asa_total_subspacesrP   rQ   ÚmaxÚminÚgetattrÚintÚasa_target_subspaces)r   Ústepr   ÚtotalÚ	mul_coeffÚexponents         r   rX   ÚAdamssModel._schedule_thresholdÜ   sç   € à×)Ñ)×-Ñ-¬d´4¸×8LÑ8LÓ3MÈyÓ.YÐ[\Ó]ˆØ�A‹:Ø×1Ñ1Ó3ˆEØ�A‹:Øà×$Ñ$Ó$ØØ×(Ñ(Ó(Ø˜t×&8Ñ&8Ñ8¸V×=PÑ=PÐSY×SeÑSeÑ=eÑfÑfˆIÜ˜C¤ S¨)Ó!4Ó5ˆIÜ˜vÐ'>ÀÓDˆHÜ�v×2Ñ2°e×>YÑ>YÑ6YÐ^gÑ^qÑ5rÑrÓsÐsà×.Ñ.Ð.r!   c                óþ   • SnU R                   R                  5        H\  n[        U[        5      (       d  M  UR                  (       d  M-  U[        [        UR                  R                  5       5      5      -  nM^     U$ )zAGet total number of subspaces from model (sum across all layers).r   )r   rR   r*   r   r-   rf   rg   Úvalues)r   ro   rA   s      r   rh   Ú$AdamssModel._get_asa_total_subspacesî   s`   € àˆØ—j‘j×(Ñ(Ö*ˆFÜ˜&¤+×.Ó.°6×3G×3GÑ3GØœœd 6×#7Ñ#7×#>Ñ#>Ó#@ÓAÓBÑB’ñ +ð ˆr!   c                ó¢  • / nU H»  nUR                   R                  US5      nUR                  U   nUR                  U   nUR                  U   n	UR
                  U   n
[        U5       HQ  nX{   b  X‹   c  M  X{   X›   -  R                  5       X‹   X«   -  R                  5       -   nUR                  X[U45        MS     M½     U(       d  g[        R                  " U Vs/ s H  oÝS   PM	     sn5      nU[        U5      :¼  a  [        S5      nO*[        R                  " U* U5      S   R                  5       * nU H‹  u  nnn[        U5      U:„  nUR                  U   nUR                   U   nU[        U5      :  a  UUU   l        U(       d
  SUU   l        U[        U5      :  d  Mn  UUU   l        U(       a  M�  SUU   l        M�     gs  snf )z¼
Apply **global** top-K masking across all layers.

Collects importance scores from every subspace in every layer, ranks them globally, and keeps only the top
``target_subspaces`` active.
r   Né   z-inf)r-   r@   rS   Úexp_avg_ipt_BÚexp_avg_unc_AÚexp_avg_unc_BÚrangeÚmeanÚappendrI   ÚstackÚlenÚfloatÚkthvalueÚitemÚadamss_AÚadamss_BÚrequires_gradÚgrad)r   Útarget_subspacesr   r^   Úsubspace_scoresrA   Ún_subspacesÚipt_AÚipt_BÚunc_AÚunc_BÚiÚscoreÚsÚ
all_scoresÚ	thresholdÚidxÚ	is_activeÚparam_AÚparam_Bs                       r   rY   Ú"AdamssModel._global_mask_to_targetö   sÃ  € ð (*ˆÛ ˆFØ ×.Ñ.×2Ñ2°<ÀÓCˆKØ×(Ñ(¨Ñ6ˆEØ×(Ñ(¨Ñ6ˆEØ×(Ñ(¨Ñ6ˆEØ×(Ñ(¨Ñ6ˆEä˜;Ö'�Ø‘8Ñ# u¡xÑ'7ÙØ™ E¡HÑ,×2Ñ2Ó4¸¹À5Á8Ñ8K×7QÑ7QÓ7SÑS�Ø×&Ñ&¨°5Ð'9Ö:ó	 (ñ !ö Øô —[’[±Ó!@²¨1 A¤$±Ñ!@ÓAˆ
Øœs ?Ó3Ó3Ü˜f›‰IäŸš¨¨Ð5EÓFÀqÑI×NÑNÓPÐPˆIó #2ÑˆF�C˜Ü˜e› yÑ0ˆIØ—o‘o lÑ3ˆGØ—o‘o lÑ3ˆGØ”S˜“\Ó!Ø-6�˜‘Ô*Þ Ø(,�G˜C‘LÔ%Ø”S˜“\Õ!Ø-6�˜‘Ô*ß �yØ(,�G˜C‘LÖ%ò #2ùò "As   Ã G)r   )FN)r   Úboolr   zOptional[dict]ÚreturnÚNone)r#   r   r   r   r6   ú	nn.Moduler7   r   r8   r›   r9   r   )r   r   r#   r   rA   r   r™   rš   )r#   r   r   r   r6   r›   r™   r›   )r[   rl   r™   rš   )rn   rl   r™   zOptional[int])r™   rl   )r‡   rl   r   r   r^   Úlistr™   rš   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Ú__annotations__r   Útuner_layer_clsr	   Útarget_module_mappingr   Ústaticmethodr%   r<   r4   r3   ra   rX   rh   rY   Ú__static_attributes__Ú__classcell__)r   s   @r   r   r      s÷   ø‡ ñ ðD €FˆCÓØ"�n€OØPÐð jnðrØ>BðrØXfðrà	÷rð rð ñ>ó ð>ð'Jà#ð'Jð ð'Jð ð	'Jð
 ð'Jð ð'Jð ô'JôRSð!à#ð!ð ð!ð ð	!ð 
ô!ôF2:ôh/ô$÷,-ò ,-r!   r   )Ú
__future__r   Útypingr   rI   r   Úpeft.tuners.tuners_utilsr   r   r   Ú
peft.utilsr	   r   r   Úlayerr   r   r   © r!   r   Ú<module>r®      s4   ðõ #å ã Ý ç ZÑ Zõõ !ß &ôC-�)õ C-r!   