ó
    >: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)Úgather_params_ctxé   )ÚPromptTuningInitc                   ó2   ^ • \ rS rSrSrU 4S jrS rSrU =r$ )ÚPromptEmbeddingé   aV  
The model to encode virtual tokens into prompt embeddings.

Args:
    config ([`PromptTuningConfig`]): The configuration of the prompt embedding.
    word_embeddings (`torch.nn.Module`): The word embeddings of the base transformer model.

**Attributes**:
    - **embedding** (`torch.nn.Embedding`) -- The embedding layer of the prompt embedding.

Example:

```py
>>> from peft import PromptEmbedding, PromptTuningConfig

>>> config = PromptTuningConfig(
...     peft_type="PROMPT_TUNING",
...     task_type="SEQ_2_SEQ_LM",
...     num_virtual_tokens=20,
...     token_dim=768,
...     num_transformer_submodules=1,
...     num_attention_heads=12,
...     num_layers=12,
...     prompt_tuning_init="TEXT",
...     prompt_tuning_init_text="Predict if sentiment of this review is positive, negative or neutral",
...     tokenizer_name_or_path="t5-base",
... )

>>> # t5_model.shared is the word embeddings of the base model
>>> prompt_embedding = PromptEmbedding(config, t5_model.shared)
```

Input Shape: (`batch_size`, `total_virtual_tokens`)

Output Shape: (`batch_size`, `total_virtual_tokens`, `token_dim`)
c                 ó\  >• [         TU ]  5         UR                  UR                  -  n[        R
                  R                  X1R                  5      U l        UR                  [        R                  :X  aù  UR                  (       dè  UR                  n[        R                  " SXC4[        R                  S9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
                  R1                  U5      U R                  l        g UR                  [        R2                  :X  Ga}  UR                  (       Gdj  SSKJn  UR8                  =(       d    0 nUR;                  SS 5        UR<                  " UR>                  40 UD6n	UR@                  n
U	" U
5      S   n[C        U5      nX³:”  a  US U nO!X³:  a  [D        RF                  " X;-  5      nX\-  nUS U n[        RH                  " U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
                  R1                  U5      U R                  l        g g g ! , (       d  f       GNú= f! , (       d  f       Np= f)Nr   )Údtype)ÚAutoTokenizerÚtrust_remote_codeÚ	input_ids)%ÚsuperÚ__init__Únum_virtual_tokensÚnum_transformer_submodulesÚtorchÚnnÚ	EmbeddingÚ	token_dimÚ	embeddingÚprompt_tuning_initr   ÚSAMPLE_VOCABÚinference_modeÚnum_embeddingsÚrandintÚlongÚtoÚweightÚdevicer   Ú
parametersÚdetachÚcloneÚfloat32Ú	ParameterÚTEXTÚtransformersr   Útokenizer_kwargsÚpopÚfrom_pretrainedÚtokenizer_name_or_pathÚprompt_tuning_init_textÚlenÚmathÚceilÚ
LongTensor)ÚselfÚconfigÚword_embeddingsÚtotal_virtual_tokensÚ
vocab_sizeÚinit_token_idsÚword_embedding_weightsr   r'   Ú	tokenizerÚ	init_textÚnum_text_tokensÚnum_repsÚ	__class__s                €Ú\/home/mande/repo/quber/.venv/lib/python3.13/site-packages/peft/tuners/prompt_tuning/model.pyr   ÚPromptEmbedding.__init__>   sz  ø€ Ü‰ÑÔà%×8Ñ8¸6×;\Ñ;\Ñ\ÐÜŸ™×+Ñ+Ð,@×BRÑBRÓSˆŒØ×$Ñ$Ô(8×(EÑ(EÓEÈf×Nc×Ncà(×7Ñ7ˆJÜ"Ÿ]š]¨1¨jÐ:QÔY^×YcÑYcÑd×gÑgØ×&Ñ&×-Ñ-óˆNô # ?×#=Ñ#=Ó#?Õ@Ù)8¸Ó)H×)OÑ)OÓ)Q×)WÑ)WÓ)YÐ&÷ Aà%;×%>Ñ%>¼u¿}¹}Ó%MÐ"Ü$)§H¡H×$6Ñ$6Ð7MÓ$NˆD�N‰NÕ!à×&Ñ&Ô*:×*?Ñ*?Ô?È×H]×H]ÐH]Ý2à%×6Ñ6×<¸"Ðð × Ñ Ð!4°dÔ;Ø%×5Ò5°f×6SÑ6SÑhÐWgÑhˆIØ×6Ñ6ˆIÙ& yÓ1°+Ñ>ˆNä! .Ó1ˆOØÓ5Ø!/Ð0EÐ1EÐ!F‘Ø Ó7ÜŸ9š9Ð%9Ñ%KÓL�Ø!/Ñ!:�Ø+Ð,AÐ-AÐBˆNÜ"×-Ò-¨nÓ=×@Ñ@À×AWÑAW×A^ÑA^Ó_ˆNÜ" ?×#=Ñ#=Ó#?Õ@Ù)8¸Ó)H×)OÑ)OÓ)Q×)WÑ)WÓ)YÐ&÷ Aà%;×%>Ñ%>¼u¿}¹}Ó%MÐ"Ü$)§H¡H×$6Ñ$6Ð7MÓ$NˆD�N‰NÕ!ð- I^Ð?÷ AÖ@ú÷0 AÕ@ús   Ã3%LÊ%LÌ
LÌ
L+c                 ó(   • U R                  U5      nU$ )N©r   )r0   ÚindicesÚprompt_embeddingss      r<   ÚforwardÚPromptEmbedding.forwardf   s   € à ŸN™N¨7Ó3ÐØ Ð ó    r?   )	Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   rB   Ú__static_attributes__Ú__classcell__)r;   s   @r<   r   r      s   ø† ñ#õJ&O÷P!ð !rD   r   )	r-   r   Úpeft.utils.integrationsr   r1   r   r   ÚModuler   © rD   r<   Ú<module>rO      s)   ðó ã å 5å $ôQ!�e—h‘h—o‘oõ Q!rD   