ó
    >:jÚ  ã                  óè   • S SK Jr  S SKrS SKJrJr  S SKrS SKJr  SSK	J
r
  \\R                  \R                  \R                  4   r " S S\R                  5      r " S S	\R                  5      rg)
é    )ÚannotationsN)ÚOptionalÚUnioné   )ÚXLoraConfigc                  ó2   ^ • \ rS rSrSU 4S jjrS rSrU =r$ )ÚTemperatureScaledSoftmaxé   c                ó`   >• [         TU ]  5         Xl        [        R                  " SS9U l        g )Néÿÿÿÿ)Údim)ÚsuperÚ__init__ÚtemperatureÚnnÚSoftmaxÚsoftmax)Úselfr   Ú	__class__s     €ÚY/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/tuners/xlora/classifier.pyr   Ú!TemperatureScaledSoftmax.__init__   s$   ø€ Ü‰ÑÔØ&ÔÜ—z’z bÑ)ˆ�ó    c                ó@   • XR                   -  nU R                  U5      $ )N)r   r   )r   ÚlogitsÚscaled_logitss      r   ÚforwardÚ TemperatureScaledSoftmax.forward"   s   € à×!1Ñ!1Ñ1ˆà�|‰|˜MÓ*Ð*r   )r   r   )g      ð?)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__r   r   Ú__static_attributes__Ú__classcell__©r   s   @r   r	   r	      s   ø† ÷*÷
+ð +r   r	   c                  ó�   ^ • \ rS rSrSr          S	U 4S jjr  S
     SS jjr  S
     SS jjrSS jrSS jr	Sr
U =r$ )ÚXLoraClassifieré)   z/
A classifier to select LoRA layers for XLora.
c           	     ó®  >• [         T
U ]  5         X0l        X@l        X l        / U l        [        U R                  R                  S9U l        UR                  U l
        SU l        [        UR                  5       5      R                  U l        UR                  S:„  n/ nU R                  R                   S:X  a§  UR"                  (       aL  [$        R&                  " UR(                  X4-  SS9R+                  U5      R+                  U R                  5      nGO¦[$        R&                  " UR(                  USS9R+                  U5      R+                  U R                  5      nGO\U R                  R                   S::  a  [-        S5      eUR/                  [$        R&                  " UR(                  UR0                  SS9R+                  U5      R+                  U R                  5      5        UR/                  [$        R2                  " 5       5        U(       a-  UR/                  [$        R4                  " UR                  S	95        [7        UR                   S
-
  5       H¾  n	UR/                  [$        R&                  " UR0                  UR0                  SS9R+                  U5      R+                  U R                  5      5        UR/                  [$        R2                  " 5       5        U(       d  M‘  UR/                  [$        R4                  " UR                  S	95        MÀ     UR"                  (       aK  [$        R&                  " UR0                  X4-  SS9R+                  U5      R+                  U R                  5      nOH[$        R&                  " UR0                  USS9R+                  U5      R+                  U R                  5      n[$        R8                  " / UQUP76 U l        g)z¡
Construct an X-LoRA classifier from a model, config and some metadata. Note that n_layers is the number of LoRA
adapter layers, not the number of model layers.
)r   Fg        r   T)Úbiasr   z'X-LoRA depth must be strictly positive.)Úpé   N)r   r   Ú	n_classesÚn_layersÚconfigÚlog_scalingsr	   Úsoftmax_temperaturer   Úscaling_pass_valueÚoverride_scaling_pass_valueÚscalings_loggingÚnextÚ
parametersÚdtypeÚxlora_dropout_pÚxlora_depthÚlayerwise_scalingsr   ÚLinearÚhidden_sizeÚtoÚ
ValueErrorÚappendÚ
xlora_sizeÚReLUÚDropoutÚrangeÚ
SequentialÚlayers)r   Úmodelr.   r,   r-   ÚdeviceÚadd_dropoutrD   ÚlastÚ_r   s             €r   r   ÚXLoraClassifier.__init__.   sÀ  ø€ ô 	‰ÑÔà"ŒØ ŒØŒØˆÔÜ/¸D¿K¹K×<[Ñ<[Ñ\ˆŒØ39×3LÑ3LˆÔ(à %ˆÔä˜%×*Ñ*Ó,Ó-×3Ñ3ˆŒ
Ø×,Ñ,¨sÑ2ˆàˆØ�;‰;×"Ñ" aÓ'Ø×(×(Ü—y’y ×!3Ñ!3°YÑ5IÐPTÑU×XÑXÐY_Ó`×cÑcÐdh×dnÑdnÓo’ä—y’y ×!3Ñ!3°YÀTÑJ×MÑMÈfÓU×XÑXÐY]×YcÑYcÓd’à�{‰{×&Ñ&¨!Ó+Ü Ð!JÓKÐKà�M‰Mœ"Ÿ)š) F×$6Ñ$6¸×8IÑ8IÐPTÑU×XÑXÐY_Ó`×cÑcÐdh×dnÑdnÓoÔpà�M‰Mœ"Ÿ'š'›)Ô$ÞØ—‘œbŸjšj¨6×+AÑ+AÑBÔCä˜6×-Ñ-°Ñ1Ö2�Ø—‘œbŸiši¨×(9Ñ(9¸6×;LÑ;LÐSWÑX×[Ñ[Ð\bÓc×fÑfÐgk×gqÑgqÓrÔsà—‘œbŸgšg›iÔ(ß�;Ø—M‘M¤"§*¢*¨v×/EÑ/EÑ"FÖGñ 3ð ×(×(Ü—y’y ×!2Ñ!2°IÑ4HÈtÑT×WÑWÐX^Ó_×bÑbÐcg×cmÑcmÓn‘ä—y’y ×!2Ñ!2°IÀDÑI×LÑLÈVÓT×WÑWÐX\×XbÑXbÓc�Ü—m’mÐ2 VÐ2¨TÒ2ˆ�r   c                óP  • Ub+  UR                   S   nUR                  nUR                   S   nO*UR                   S   nUR                  nUR                   S   n[        R                  " XWU R                  U R
                  4U R                  5      R                  X`R                  S9$ )a0  
Make some dummy scalings for the scalings pass (the one to get the logits for the X-LoRA classifier). These are
of shape (batch_size, seq_len, n_layers, n_classes) and filled with the override scalings pass value. Note that
n_layers is the number of LoRA adapter layers, not the number of model layers.
r   r   )rF   r6   )	ÚshaperF   ÚtorchÚfullr-   r,   r2   r<   r6   )r   Ú	input_idsÚinputs_embedsÚargsÚkwargsÚ
batch_sizerF   Úseq_lens           r   Úmake_dummy_scalingsÚ#XLoraClassifier.make_dummy_scalingse   sš   € ð Ñ Ø"Ÿ™¨Ñ+ˆJØ×%Ñ%ˆFØ—o‘o aÑ(‰Gà&×,Ñ,¨QÑ/ˆJØ"×)Ñ)ˆFØ#×)Ñ)¨!Ñ,ˆGä�zŠzØ $§-¡-°·±Ð@Ø×,Ñ,ó
÷ ‰"�F§*¡*ˆ"Ð
-ð	.r   c                óp  • Ub  UR                   S   nUR                   S   nOUR                   S   nUR                   S   nUR                  nUS   n	U R                  R                  U	5      n
U R                  R
                  (       d/  U
R                  S5      n
U
R                  SSU R                  S5      n
U
R                  XgU R                  U R                  5      nU R                  R                  (       a  U R                  U5      nU R                  (       a  U R                  R                  U5        U$ )zd
Using the hidden states of the model, predict `n_classes` LoRA alpha values. Returns the scalings.
r   r   r   r+   )rL   Úhidden_statesrD   r   r.   r9   Ú	unsqueezeÚexpandr-   Úreshaper,   Úenable_softmaxr   r3   r/   r>   )r   ÚresultrO   rP   rQ   rR   rS   rT   rX   Úhidden_stater   Úscalingss               r   r   ÚXLoraClassifier.forward   s  € ð Ñ Ø"Ÿ™¨Ñ+ˆJØ—o‘o aÑ(‰Gà&×,Ñ,¨QÑ/ˆJØ#×)Ñ)¨!Ñ,ˆGà×,Ñ,ˆà$ RÑ(ˆð —‘×$Ñ$ \Ó2ˆð
 �{‰{×-×-Ø×%Ñ% aÓ(ˆFØ—]‘] 2 r¨4¯=©=¸"Ó=ˆFð —>‘> *°t·}±}ÀdÇnÁnÓUˆð �;‰;×%×%Ø—|‘| HÓ-ˆHà× × Ø×Ñ×$Ñ$ XÔ.àˆr   c                óÚ   • 0 n[        U R                  5       HO  u  p#UR                  S   nXA;  a
  U/U/4X'   M#  X   S   R                  U5        X   S   R                  U5        MQ     U$ )a,  
Returns bucketed scalings, bucketed by seq_len. Each value consists of the positions (the first) and the
associated tensors. The positions are paired with the associated tensors and give the position in the scaling
log. Each scaling is a tensor of shape (batch_size, seq_len, n_layers, n_classes)).
r   r   )Ú	enumerater/   rL   r>   )r   Úseqlens_mapÚiÚscalingrT   s        r   Ú_get_bucketed_scalingsÚ&XLoraClassifier._get_bucketed_scalings­   s{   € ð HJˆÜ# D×$5Ñ$5Ö6‰JˆAØ—m‘m AÑ&ˆGØÓ)Ø)*¨¨g¨YÐ'7�Ó$àÑ$ QÑ'×.Ñ.¨qÔ1ØÑ$ QÑ'×.Ñ.¨wÖ7ñ 7ð Ðr   c                óv   • Uc  SU R                   -  U l        OXl        U R                  U R                  l        g )Nr   )r,   r2   r.   r1   )r   Úvalues     r   Ú _set_override_scaling_pass_valueÚ0XLoraClassifier._set_override_scaling_pass_value¾   s0   € Ø‰=Ø/0°4·>±>Ñ/AˆDÕ,à/4Ô,Ø)-×)IÑ)Iˆ�‰Õ&r   )	r.   r6   rD   r/   r,   r-   r2   r3   r   )
rE   z	nn.Moduler.   r   r,   Úintr-   rl   rF   ztorch.device)NN)rO   zOptional[torch.LongTensor]rP   zOptional[torch.FloatTensor]Úreturnztorch.Tensor)rm   z/dict[int, tuple[list[int], list[torch.Tensor]]])ri   zUnion[Number, None])r   r   r    r!   Ú__doc__r   rU   r   rf   rj   r"   r#   r$   s   @r   r&   r&   )   s¤   ø† ñð53àð53ð ð53ð ð	53ð
 ð53ð ÷53ðr 15Ø59ð.à-ð.ð 3ð.ð 
õ.ð: 15Ø59ð	,ð .ð,ð 3ð	,ð 
õ,ô\÷"Jò Jr   r&   )Ú
__future__r   ÚbuiltinsÚtypingr   r   rM   Útorch.nnr   r.   r   rl   ÚfloatÚboolÚNumberÚModuler	   r&   © r   r   Ú<module>rx      s\   ðõ #ã ß "ã Ý å ð 
ˆx�|‰|˜XŸ^™^¨X¯]©]Ð:Ñ	;€ô
+˜rŸy™yô 
+ôZJ�b—i‘iõ ZJr   