ó
    >:j!  ã                   ój   • S SK r S SKrS SKJr  S SKJr   " S S\R                  R                  5      rg)é    N)ÚCrossEntropyLoss)Úgather_params_ctxc                   óT   ^ • \ rS rSrSrU 4S jrS rS rS rS r	\
S 5       rS	rU =r$ )
ÚCPTEmbeddingé   zÇ
CPTEmbedding is a custom embedding layer designed for Context-aware Prompt Tuning (CPT) in PEFT. It initializes
embeddings, applies prompt-specific projections, and computes loss using label masks.
c                 óz  >• [         TU ]  5         [        R                  " U5      U l        UR
                  n[        R                  R                  X1R                  5      U l
        UR                  (       dû  UR
                  [        UR                  5      :X  d   e[        R                  " UR                  5      R                  UR                   R"                  5      n[%        UR'                  5       5         U" U5      R)                  5       R+                  5       nSSS5        WR                  [        R,                  5      n[        R                  R/                  U5      U R                  l        U R                  R1                  S5        [        R                  R                  X1R                  5      U l        [        R4                  " U R2                  R                   5      R                  [        R,                  5      U R2                  R                   l        U R9                  5         g! , (       d  f       GN= f)a  
Initializes the CPTEmbedding module.

Args:
    config (Namespace):
        Configuration object containing model hyperparameters and CPT-specific settings.
    word_embeddings (torch.nn.Embedding):
        The base word embedding layer used to initialize CPT embeddings.
NF)ÚsuperÚ__init__ÚcopyÚdeepcopyÚconfigÚnum_virtual_tokensÚtorchÚnnÚ	EmbeddingÚ	token_dimÚ	embeddingÚinference_modeÚlenÚcpt_token_idsÚ
LongTensorÚtoÚweightÚdevicer   Ú
parametersÚdetachÚcloneÚfloat32Ú	ParameterÚrequires_grad_Údelta_embeddingÚ
zeros_likeÚdataÚset_updated_tokens)Úselfr   Úword_embeddingsr   Úinit_token_idsÚword_embedding_weightsÚ	__class__s         €ÚR/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/tuners/cpt/model.pyr
   ÚCPTEmbedding.__init__   sŽ  ø€ ô 	‰ÑÔÜ—m’m FÓ+ˆŒØ#×6Ñ6Ðô Ÿ™×+Ñ+Ð,>×@PÑ@PÓQˆŒð ×$×$Ø×,Ñ,´°F×4HÑ4HÓ0IÓIÐIÐIä"×-Ò-¨f×.BÑ.BÓC×FÑFÀ×G]ÑG]×GdÑGdÓeˆNÜ" ?×#=Ñ#=Ó#?Õ@Ù)8¸Ó)H×)OÑ)OÓ)Q×)WÑ)WÓ)YÐ&÷ Aà%;×%>Ñ%>¼u¿}¹}Ó%MÐ"Ü$)§H¡H×$6Ñ$6Ð7MÓ$NˆD�N‰NÔ!à�‰×%Ñ% eÔ,ô  %Ÿx™x×1Ñ1Ð2D×FVÑFVÓWˆÔÜ+0×+;Ò+;¸D×<PÑ<P×<WÑ<WÓ+X×+[Ñ+[Ô\a×\iÑ\iÓ+jˆ×Ñ×#Ñ#Ô(ð 	×ÑÕ!÷ AÖ@ús   Ã7%H+È+
H:c                 ó   • [         R                  " 5          U R                  U5      nSSS5        U R                  5       U R                  R
                  l        U R	                  U5      nWU-   $ ! , (       d  f       NM= f)zã
Computes the prompt embeddings and applies delta adjustments.

Args:
    indices (torch.Tensor):
        Indices of the tokens to be embedded.

Returns:
    torch.Tensor:
        Sum of prompt embeddings and delta embeddings.
N)r   Úno_gradr   Úget_projectionr!   r   r#   )r%   ÚindicesÚprompt_embeddingsÚdelta_prompt_embeddingss       r*   ÚforwardÚCPTEmbedding.forwardA   se   € ô �]Š]�_Ø $§¡¨wÓ 7Ð÷ ð ,0×+>Ñ+>Ó+@ˆ×Ñ×#Ñ#Ô(à"&×"6Ñ"6°wÓ"?Ðà Ð#:Ñ:Ð:÷ �_ús   –A/Á/
A=c                 óš  ^• [         R                  " U R                  R                  5      R	                  5       n[         R
                  " US5      S:H  n[         R
                  " US5      S:H  n[         R
                  " US5      S:H  nX#-  U-  mTR                  SS5      mU4S jnU R                  R                  R                  U5        g)za
Sets up a backward hook to selectively update token gradients based on the CPT token type mask.
é   é   é   é   éÿÿÿÿc                 óD   >• U TR                  U R                  5      -  n U $ )N)r   r   )ÚgradÚmasks    €r*   Úbackward_hookÚ6CPTEmbedding.set_updated_tokens.<locals>.backward_hooka   s   ø€ Ø˜$Ÿ'™' $§+¡+Ó.Ñ.ˆDØˆKó    N)
r   ÚTensorr   Úcpt_tokens_type_maskÚlongÚ	remainderÚviewr!   r   Úregister_hook)r%   Útensor_ICL_maskÚmask_input_templateÚ
mask_inputÚmask_output_templater=   r<   s         @r*   r$   ÚCPTEmbedding.set_updated_tokensV   s¦   ø€ ô  Ÿ,š, t§{¡{×'GÑ'GÓH×MÑMÓOˆÜ#Ÿošo¨o¸qÓAÀQÑFÐÜ—_’_ _°aÓ8¸AÑ=ˆ
Ü$Ÿš¨ÀÓBÀaÑGÐØ"Ñ/Ð2FÑFˆØ�y‰y˜˜QÓˆõ	ð 	×Ñ×#Ñ#×1Ñ1°-Õ@r?   c                 óB  • U R                   R                  nSnU R                   R                  [        R                  " [        R
                  " U R                   R                  S-  /5      5      -  nU R                   R                  [        R                  " [        R
                  " U R                   R                  S-  /5      5      -  n[        R                  " [        R
                  " U5      5      R                  [        R                  5      U-  n[        R
                  " U5      R                  5       nX5US:„  [        R                  " US5      S:H  -  '   X5US:„  [        R                  " US5      S:H  -  '   XEUS:„  [        R                  " US5      S:H  -  '   U$ )Ng»½×Ùß|Û=i   r   r5   r6   r8   r7   )r   rA   Úopt_projection_format_epsilonr   Úsqrtr@   r   Úopt_projection_epsilonÚ	ones_liker   r   rB   rC   )r%   rA   Ú	MIN_VALUEÚnormalized_format_epsÚnormalized_input_epsÚepsilons         r*   Úget_epsilonÚCPTEmbedding.get_epsilong   sX  € Ø#Ÿ{™{×?Ñ?Ðàˆ	ð !%§¡× IÑ IÌEÏJÊJÜ�LŠL˜$Ÿ+™+×/Ñ/°$Ñ6Ð7Ó8óM
ñ !
Ðð  $Ÿ{™{×AÑAÄEÇJÂJÜ�LŠL˜$Ÿ+™+×/Ñ/°$Ñ6Ð7Ó8óE
ñ  
Ðô —/’/¤%§,¢,Ð/CÓ"DÓE×HÑHÌÏÉÓWÐZcÑcˆÜ$Ÿ|š|Ð,@ÓA×FÑFÓHÐà`uÐ%¨Ñ)¬e¯oªoÐ>RÐTUÓ.VÐZ[Ñ.[Ñ\Ñ]Ø`uÐ%¨Ñ)¬e¯oªoÐ>RÐTUÓ.VÐZ[Ñ.[Ñ\Ñ]Ø`tÐ%¨Ñ)¬e¯oªoÐ>RÐTUÓ.VÐZ[Ñ.[Ñ\Ñ]àˆr?   c           	      óR  • [         R                  " 5          U R                  R                  R	                  5       R                  U R                  R                  R                  5      n[         R                  " USSS9nUS:„  n[         R                  " U5      (       ao  U R                  5       R                  U R                  R                  R                  5      nX==   XC   X#   R                  XC   S9-  R                  SS5      -  ss'   UsSSS5        $ ! , (       d  f       g= f)zQ
Applies epsilon-based projection to the delta embeddings to control their norm.
r7   r6   )ÚpÚdimr   )Úminr9   N)r   r-   r!   r   r   r   r   ÚnormÚanyrT   ÚclamprD   )r%   Únew_embeddings_weightsÚ
token_normÚprojection_maskrS   s        r*   r.   ÚCPTEmbedding.get_projection}   sã   € ô �]Š]�_Ø%)×%9Ñ%9×%@Ñ%@×%FÑ%FÓ%H×%KÑ%KÈD×L`ÑL`×LgÑLg×LnÑLnÓ%oÐ"ÜŸšÐ$:¸aÀQÑGˆJà(¨1™nˆOÜ�yŠy˜×)Ñ)Ø×*Ñ*Ó,×/Ñ/°×0DÑ0D×0KÑ0K×0RÑ0RÓS�Ø&Ó7ØÑ,°
Ñ0K×0QÑ0QÐV]ÑVnÐ0QÐ0oÑpß‘$�r˜1“+ñÓ7ð *÷ �_�_ús   –C8DÄ
D&c                 ó,  • U R                   R                  nU R                   nUR                  U5      nUSSS2SS24   R                  5       nUSSS24   R                  5       nUSSS24   R                  5       nUR	                  5       R                  5       S:g  R                  5       n	UR                  u  p«n[        SSS9nU" UR                  X«-  U5      UR                  X«-  5      5      nUR                  X«5      nU	R	                  5       R                  5       R                  5       n[        U
5       H»  nUU   S:„  UU   S	-  S:H  -  nUU   U   R                  5       n[        R                  " UU   5      R                  US
9R                  5       nSn[        R                  " US/5       H  nUUUU   U:H  '   UUR                   -  nM     UR"                  S:X  d  M®  UU==   U-  ss'   M½     Xé   Xù   -  R%                  5       nXàl        U $ )aü  
Computes the loss for CPT models with optional exponential decay.

Args:
    base_model_output (ModelOutput):
        Output from the base model containing logits.
    labels (torch.Tensor):
        Ground-truth labels for the input tokens.
    cpt_type_mask (torch.Tensor):
        Token type mask used for filtering valid loss terms.
    config (Namespace):
        Configuration object containing loss-related hyperparameters.

Returns:
    ModelOutput:
        The base model output with computed loss.
.Nr9   r6   iœÿÿÿÚnone)Ú	reductionÚignore_indexr   r5   )r   Údecay)Úlogitsr   r   Ú
contiguousr   r   ÚboolÚshaper   rD   ÚfloatÚrangeÚuniquer   rO   ÚflipÚopt_loss_decay_factorÚopt_weighted_loss_typeÚmeanÚloss)Úbase_model_outputÚlabelsÚcpt_type_maskr   r   Ú	lm_logitsÚshift_logitsÚshift_labelsÚshift_cpt_type_maskÚshift_labels_boolÚ
batch_sizeÚ
seq_lengthÚ
vocab_sizeÚloss_fctrq   Úshift_labels_weightsÚiÚ
idx_labelsÚ
labels_idsÚexponential_decayÚdecay_valueÚlabel_mask_idxs                         r*   Úcalculate_lossÚCPTEmbedding.calculate_loss�   s'  € ð( #×)Ñ)×0Ñ0ˆà%×,Ñ,ˆ	Ø—‘˜6Ó"ˆð !  c r cª1 Ñ-×8Ñ8Ó:ˆØ˜c 1¡2˜g‘×1Ñ1Ó3ˆØ+¨C°±¨GÑ4×?Ñ?ÓAÐà)×/Ñ/Ó1×8Ñ8Ó:¸dÑB×HÑHÓJÐØ-9×-?Ñ-?Ñ*ˆ
 
ô $¨fÀ4ÑHˆÙØ×Ñ˜jÑ5°zÓBÀL×DUÑDUÐV`ÑVmÓDnó
ˆð �y‰y˜Ó0ˆà0×6Ñ6Ó8×?Ñ?ÓA×GÑGÓIÐä�zÖ"ˆAØ-¨aÑ0°1Ñ4Ð9LÈQÑ9OÐRSÑ9SÐWXÑ9XÑYˆJØ,¨QÑ/°
Ñ;×BÑBÓDˆJä %§¢Ð0CÀAÑ0FÓ G× JÑ JÐRXÐ JÐ Y× _Ñ _Ó aÐØˆKÜ"'§*¢*¨Z¸!¸Ö"=�ØNYÐ!Ð"5°aÑ"8¸NÑ"JÑKØ˜v×;Ñ;Ñ;’ñ #>ð ×,Ñ,°Õ7Ø$ QÓ'Ð+<Ñ<Õ'ñ #ð Ñ'Ð*>Ñ*QÑQ×WÑWÓYˆà!%Ôà Ð r?   )r   r!   r   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r
   r2   r$   rT   r.   Ústaticmethodr…   Ú__static_attributes__Ú__classcell__)r)   s   @r*   r   r      s7   ø† ñõ
""òH;ò*Aò"ò,*ð$ ñ:!ó ö:!r?   r   )	r   r   Útorch.nnr   Úpeft.utils.integrationsr   r   ÚModuler   © r?   r*   Ú<module>r“      s)   ðó ã Ý %å 5ôs!�5—8‘8—?‘?õ s!r?   