§
    ‚Štjö]  ã                   ó¾  — d Z ddlmZ ddlZddlZddlmZ ddlmZm	Z	m
Z
 ddlmZ ddlmZmZ dd	lmZ dd
lmZ ddlmZmZmZ ddlmZmZ ddlmZ ddlmZm Z m!Z!m"Z" ddl#m$Z$ ddl%m&Z& ddl'm(Z(  e"j)        e*¦  «        Z+d„ Z,d„ Z-	 	 d0dej.        dej/        dej/        dej/        dej/        dz  de0dz  de0dee         fd„Z1 G d„ d ej.        ¦  «        Z2d!„ Z3 G d"„ d#ej.        ¦  «        Z4e  G d$„ d%e¦  «        ¦   «         Z5e  G d&„ d'e5¦  «        ¦   «         Z6 e d(¬)¦  «         G d*„ d+e5e¦  «        ¦   «         Z7 e d,¬)¦  «         G d-„ d.e5¦  «        ¦   «         Z8g d/¢Z9dS )1zPyTorch CTRL model.é    )ÚCallableN)Únn)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELossé   )Úinitialization)ÚCacheÚDynamicCache)ÚGenerationMixin)Úcreate_causal_mask)ÚBaseModelOutputWithPastÚCausalLMOutputWithPastÚSequenceClassifierOutput)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)ÚTransformersKwargsÚauto_docstringÚcan_return_tupleÚlogging)Úmerge_with_config_defaults)Úcapture_outputsé   )Ú
CTRLConfigc                 óN   — dt          j        dd|dz  z  |z  ¦  «        z  }| |z  S )Nr   i'  é   )ÚtorchÚpow)ÚposÚiÚd_model_sizeÚangle_ratess       úd/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/ctrl/modeling_ctrl.pyÚ
angle_defnr%   -   s0   € Ø•e”i ¨¨Q°!©V©¸Ñ'DÑEÔEÑE€KØ�ÑÐó    c                 óì  — t          t          j        | t          j        ¬¦  «                             |¦  «                             d¦  «        t          j        |t          j        ¬¦  «                             |¦  «                             d¦  «        |¦  «        }t          j        |d d …dd d…f         ¦  «        }t          j        |d d …dd d…f         ¦  «        }t          j        ||gd¬¦  «        }|S )N)Údtyper   r   r   éÿÿÿÿ©Údim)	r%   r   ÚarangeÚint64ÚtoÚ	unsqueezeÚsinÚcosÚcat)Úpositionr"   r(   Ú
angle_radsÚsinesÚcosinesÚpos_encodings          r$   Úpositional_encodingr8   2   sÙ   € åÝŒ�X¥U¤[Ð1Ñ1Ô1×4Ò4°UÑ;Ô;×EÒEÀaÑHÔHÝŒ�\­¬Ð5Ñ5Ô5×8Ò8¸Ñ?Ô?×IÒIÈ!ÑLÔLØñô €Jõ ŒI�j    A D q D Ô)Ñ*Ô*€EÝŒi˜
 1 1 1 a d¨ d 7Ô+Ñ,Ô,€Gå”9˜e WÐ-°2Ð6Ñ6Ô6€LØÐr&   ç        ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingÚdropoutÚkwargsc                 ó®  — |€|                      d¦  «        dz  }t          j        ||                     dd¦  «        ¦  «        |z  }|�||z   }t          j                             |d¬¦  «        }t          j                             ||| j        ¬¦  «        }t          j        ||¦  «        }	|	                     dd¦  «         	                    ¦   «         }	|	|fS )Nr)   ç      à¿r   r   r*   )ÚpÚtrainingr   )
Úsizer   ÚmatmulÚ	transposer   Ú
functionalÚsoftmaxr@   rE   Ú
contiguous)
r:   r;   r<   r=   r>   r?   r@   rA   Úattn_weightsÚattn_outputs
             r$   Úeager_attention_forwardrN   B   sÈ   € ð €Ø—*’*˜R‘.”. DÑ(ˆõ ”<  s§}¢}°Q¸Ñ':Ô':Ñ;Ô;¸gÑE€LàÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2Ð(Ñ>Ô>€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€Lå”,˜|¨UÑ3Ô3€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r&   c                   ó>   ‡ — e Zd Zdˆ fd„	Z	 	 ddee         fd„Zˆ xZS )ÚMultiHeadAttentionNc                 ó"  •— t          ¦   «                              ¦   «          || _        |j        | _        |j        | _        || _        d| _        t          | j        | j        z  ¦  «        | _
        | j
        dz  | _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        t          j        | j        | j        ¦  «        | _        d S )NTrC   )ÚsuperÚ__init__ÚconfigÚn_headÚ	num_headsÚn_embdr"   Ú	layer_idxÚ	is_causalÚintÚhead_dimr?   r   ÚLinearÚWqÚWkÚWvÚdense©ÚselfrT   rX   Ú	__class__s      €r$   rS   zMultiHeadAttention.__init___   sÐ   ø€ Ý‰Œ×ÒÑÔÐØˆŒØœˆŒØ"œMˆÔØ"ˆŒØˆŒå˜DÔ-°´Ñ>Ñ?Ô?ˆŒØ”} dÑ*ˆŒå”)˜DÔ-¨tÔ/@ÑAÔAˆŒÝ”)˜DÔ-¨tÔ/@ÑAÔAˆŒÝ”)˜DÔ-¨tÔ/@ÑAÔAˆŒå”Y˜tÔ0°$Ô2CÑDÔDˆŒ
ˆ
ˆ
r&   rA   c                 óÒ  — |j         d d…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }	|                      |¦  «                             |¦  «                             dd¦  «        }
|                      |¦  «                             |¦  «                             dd¦  «        }|�|                     |
|| j        ¦  «        \  }
}t          j
        | j        j        t          ¦  «        } || |	|
||fd| j        dœ|¤Ž\  }} |j        g |¢d‘R Ž                      ¦   «         }|                      |¦  «        }||fS )Nr)   r   r   r9   )r@   r?   )Úshaper[   r]   ÚviewrH   r^   r_   ÚupdaterX   r   Úget_interfacerT   Ú_attn_implementationrN   r?   ÚreshaperK   r`   )rb   ÚvÚkÚqÚ
layer_pastr>   rA   Úinput_shapeÚhidden_shapeÚquery_statesÚ
key_statesÚvalue_statesÚattention_interfacerM   rL   s                  r$   ÚforwardzMultiHeadAttention.forwardp   s~  € ð ”g˜c˜r˜c”lˆØ8˜Ð8 bÐ8¨$¬-Ð8Ð8ˆà—w’w˜q‘z”z—’ |Ñ4Ô4×>Ò>¸qÀ!ÑDÔDˆØ—W’W˜Q‘Z”Z—_’_ \Ñ2Ô2×<Ò<¸QÀÑBÔBˆ
Ø—w’w˜q‘z”z—’ |Ñ4Ô4×>Ò>¸qÀ!ÑDÔDˆàÐ!Ø'1×'8Ò'8¸À\ÐSWÔSaÑ'bÔ'bÑ$ˆJ˜å(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð %8Ð$7ØØØØØð	%
ð Ø”Lð	%
ð 	%
ð ð	%
ð 	%
Ñ!ˆ�\ð *�kÔ)Ð;¨;Ð;¸Ð;Ð;Ð;×FÒFÑHÔHˆØ—j’j Ñ-Ô-ˆØ˜LÐ(Ð(r&   ©N©NN©Ú__name__Ú
__module__Ú__qualname__rS   r   r   ru   Ú__classcell__©rc   s   @r$   rP   rP   ^   st   ø€ € € € € ðEð Eð Eð Eð Eð Eð, Øð#)ð #)ð Ð+Ô,ð#)ð #)ð #)ð #)ð #)ð #)ð #)ð #)r&   rP   c                 óœ   — t          j        t          j        | |¦  «        t          j        ¦   «         t          j        || ¦  «        ¦  «        S rv   )r   Ú
Sequentialr\   ÚReLU)r"   Údffs     r$   Úpoint_wise_feed_forward_networkr‚   –   s5   € ÝŒ=�œ <°Ñ5Ô5µr´w±y´yÅ"Ä)ÈCÐQ]ÑB^ÔB^Ñ_Ô_Ð_r&   c                   ó>   ‡ — e Zd Zdˆ fd„	Z	 	 ddee         fd„Zˆ xZS )ÚEncoderLayerNc                 óª  •— t          ¦   «                              ¦   «          t          ||¬¦  «        | _        t	          |j        |j        ¦  «        | _        t          j	        |j        d¬¦  «        | _
        t          j	        |j        d¬¦  «        | _        t          j        |j        ¦  «        | _        t          j        |j        ¦  «        | _        d S )N©rX   g�íµ ÷Æ°>©Úeps)rR   rS   rP   Úmulti_head_attentionr‚   rW   r�   Úffnr   Ú	LayerNormÚ
layernorm1Ú
layernorm2ÚDropoutÚresid_pdropÚdropout1Údropout2ra   s      €r$   rS   zEncoderLayer.__init__›   sŸ   ø€ Ý‰Œ×ÒÑÔÐå$6°vÈÐ$SÑ$SÔ$SˆÔ!Ý2°6´=À&Ä*ÑMÔMˆŒåœ, v¤}¸$Ð?Ñ?Ô?ˆŒÝœ, v¤}¸$Ð?Ñ?Ô?ˆŒåœ
 6Ô#5Ñ6Ô6ˆŒÝœ
 6Ô#5Ñ6Ô6ˆŒˆˆr&   rA   c                 ó  — |                       |¦  «        } | j        |||f||dœ|¤Ž\  }}|                      |¦  «        }||z   }|                      |¦  «        }	|                      |	¦  «        }
|                      |
¦  «        }
||
z   }	|	S )N©rn   r>   )rŒ   r‰   r�   r�   rŠ   r‘   )rb   Úxrn   r>   rA   ÚnormedrM   Ú_Úout1Úout2Ú
ffn_outputs              r$   ru   zEncoderLayer.forward§   s®   € ð —’ Ñ#Ô#ˆØ2˜Ô2ØØØð
ð "Ø)ð
ð 
ð ð
ð 
‰ˆ�Qð —m’m KÑ0Ô0ˆØ�;‰ˆà�Š˜tÑ$Ô$ˆØ—X’X˜d‘^”^ˆ
Ø—]’] :Ñ.Ô.ˆ
Ø�jÑ ˆàˆr&   rv   rw   rx   r}   s   @r$   r„   r„   š   sn   ø€ € € € € ð
7ð 
7ð 
7ð 
7ð 
7ð 
7ð Øð	ð ð
 Ð+Ô,ðð ð ð ð ð ð ð r&   r„   c                   óH   ‡ — e Zd ZU eed<   dZdZdZdZdZ	e
edœZˆ fd„Zˆ xZS )ÚCTRLPreTrainedModelrT   ÚtransformerT)Úhidden_statesÚ
attentionsc                 óü   •— t          ¦   «                              |¦  «         t          |t          ¦  «        rDt	          j        |j        t          |j        j	        |j
        t          j        ¦  «        ¦  «         d S d S rv   )rR   Ú_init_weightsÚ
isinstanceÚ	CTRLModelÚinitÚcopy_r7   r8   rT   Ún_positionsr"   r   Úfloat)rb   r:   rc   s     €r$   r    z!CTRLPreTrainedModel._init_weightsÏ   sw   ø€ Ý‰Œ×Ò˜fÑ%Ô%Ð%Ý�f�iÑ(Ô(ð 	ÝŒJØÔ#Õ%8¸¼Ô9RÐTZÔTgÕinÔitÑ%uÔ%uñô ð ð ð ð	ð 	r&   )ry   rz   r{   r   Ú__annotations__Úbase_model_prefixÚ_supports_flash_attnÚ_supports_sdpaÚ_supports_flex_attnÚ_supports_attention_backendr„   rP   Ú_can_record_outputsr    r|   r}   s   @r$   r›   r›   Â   sv   ø€ € € € € € àÐÐÑØ%ÐØÐØ€NØÐØ"&Ðà%Ø(ðð Ðð
ð ð ð ð ð ð ð ð r&   r›   c                   óþ   ‡ — e Zd Zˆ fd„Zd„ Zd„ Zeee	 	 	 	 	 	 	 dde	j
        dz  dedz  de	j        dz  de	j
        dz  d	e	j
        dz  d
e	j        dz  dedz  dee         defd„¦   «         ¦   «         ¦   «         Zˆ xZS )r¢   c                 óV  •‡— t          ¦   «                              ‰¦  «         ‰j        | _        ‰j        | _        t          j        ‰j        ‰j        ¦  «        | _	        t          j
        ‰j        ¦  «        | _        t          j        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        t          j        ‰j        ‰j        ¬¦  «        | _        |                      dt)          ‰j        | j        t,          j        ¦  «        d¬¦  «         |                      ¦   «          d S )Nc                 ó2   •— g | ]}t          ‰|¬ ¦  «        ‘ŒS )r†   )r„   )Ú.0r!   rT   s     €r$   ú
<listcomp>z&CTRLModel.__init__.<locals>.<listcomp>â   s&   ø€ ÐaÐaÐaÀa¥¨V¸qÐ AÑ AÔ AÐaÐaÐar&   r‡   r7   F)Ú
persistent)rR   rS   rW   r"   Ún_layerÚ
num_layersr   Ú	EmbeddingÚ
vocab_sizeÚwrŽ   Ú
embd_pdropr@   Ú
ModuleListÚrangeÚhr‹   Úlayer_norm_epsilonÚ	layernormÚregister_bufferr8   r¥   r   r¦   Ú	post_init©rb   rT   rc   s    `€r$   rS   zCTRLModel.__init__Ù   sü   øø€ Ý‰Œ×Ò˜Ñ Ô Ð à"œMˆÔØ œ.ˆŒå”˜fÔ/°´Ñ?Ô?ˆŒå”z &Ô"3Ñ4Ô4ˆŒÝ”ÐaÐaÐaÐaÍ5ÐQWÔQ_ÑK`ÔK`ÐaÑaÔaÑbÔbˆŒÝœ f¤m¸Ô9RÐSÑSÔSˆŒà×ÒØÕ/°Ô0BÀDÔDUÕW\ÔWbÑcÔcÐpuð 	ñ 	
ô 	
ð 	
ð
 	�ŠÑÔÐÐÐr&   c                 ó   — | j         S rv   ©r¸   )rb   s    r$   Úget_input_embeddingszCTRLModel.get_input_embeddingsì   s	   € ØŒvˆr&   c                 ó   — || _         d S rv   rÃ   )rb   Únew_embeddingss     r$   Úset_input_embeddingszCTRLModel.set_input_embeddingsï   s   € ØˆŒˆˆr&   NÚ	input_idsÚpast_key_valuesr>   Útoken_type_idsÚposition_idsÚinputs_embedsÚ	use_cacherA   Úreturnc                 ó
  — |�|n| j         j        }|�|�t          d¦  «        ‚|�|                      |¦  «        }n|€t          d¦  «        ‚|j        dd…         }	|j        d         }
|j        }|r|€t          | j         ¬¦  «        }|�|                     ¦   «         nd}|€@t          j	        ||	d         |z   t          j
        |¬¦  «        }|                     d¦  «        }|�2|                      |¦  «        }|t          j        | j        ¦  «        z  }nd}|�!|j        dk     r|                     |
d¦  «        }t#          | j         ||||¬	¦  «        }|t          j        | j        ¦  «        z  }| j                             ||j        ¬
¦  «        | _        | j        |dd…f         }||z   |z   }|                      |¦  «        }| j        D ]} ||f||dœ|¤Ž}Œ|                      |¦  «        }t1          ||r|nd¬¦  «        S )a¡  
        Example:

        ```python
        >>> from transformers import AutoTokenizer, CTRLModel
        >>> import torch

        >>> tokenizer = AutoTokenizer.from_pretrained("Salesforce/ctrl")
        >>> model = CTRLModel.from_pretrained("Salesforce/ctrl")

        >>> # CTRL was trained with control codes as the first token
        >>> inputs = tokenizer("Opinion My dog is cute", return_tensors="pt")
        >>> assert inputs["input_ids"][0, 0].item() in tokenizer.control_codes.values()

        >>> outputs = model(**inputs)

        >>> last_hidden_states = outputs.last_hidden_state
        >>> list(last_hidden_states.shape)
        [1, 5, 1280]
        ```NzDYou cannot specify both input_ids and inputs_embeds at the same timez5You have to specify either input_ids or inputs_embedsr)   r   )rT   )r(   Údeviceé   )rT   rÌ   r>   rÉ   rË   ©rÐ   r(   r“   )Úlast_hidden_staterÉ   )rT   rÍ   Ú
ValueErrorr¸   re   rÐ   r   Úget_seq_lengthr   r,   Úlongr/   ÚnpÚsqrtr"   Úndimrf   r   r7   r.   r(   r@   r¼   r¾   r   )rb   rÈ   rÉ   r>   rÊ   rË   rÌ   rÍ   rA   ro   Ú
batch_sizerÐ   Úpast_lengthÚtoken_type_embedsÚcausal_maskÚ
pos_embedsr�   r¼   s                     r$   ru   zCTRLModel.forwardò   sr  € ðD "+Ð!6�I�I¸D¼KÔ<Qˆ	àÐ  ]Ð%>ÝÐcÑdÔdÐdØÐ Ø ŸFšF 9Ñ-Ô-ˆMˆMØÐ"ÝÐTÑUÔUÐUà#Ô)¨#¨2¨#Ô.ˆØ"Ô(¨Ô+ˆ
ØÔ%ˆàð 	?˜Ð0Ý*°$´+Ð>Ñ>Ô>ˆOà:IÐ:U�o×4Ò4Ñ6Ô6Ð6Ð[\ˆØÐÝ œ<¨°[À´_À{Ñ5RÕZ_ÔZdÐmsÐtÑtÔtˆLØ'×1Ò1°!Ñ4Ô4ˆLàÐ%Ø $§¢ ~Ñ 6Ô 6ÐØ¥¤¨Ô):Ñ!;Ô!;Ñ;ÐÐà !ÐàÐ%¨.Ô*=ÀÒ*AÐ*AØ+×0Ò0°¸RÑ@Ô@ˆNå(Ø”;Ø'Ø)Ø+Ø%ð
ñ 
ô 
ˆð 	�œ Ô!2Ñ3Ô3Ñ3ˆð !Ô-×0Ò0¸ÀmÔFYÐ0ÑZÔZˆÔØÔ& |°Q°Q°Q Ô7ˆ
à%¨
Ñ2Ð5FÑFˆàŸš ]Ñ3Ô3ˆà”ð 	ð 	ˆAØ˜AØðà*Ø*ðð ð ð	ð ˆMˆMð Ÿš }Ñ5Ô5ˆå&Ø+Ø/8ÐB˜O˜O¸dð
ñ 
ô 
ð 	
r&   )NNNNNNN)ry   rz   r{   rS   rÄ   rÇ   r   r   r   r   Ú
LongTensorr
   ÚFloatTensorÚboolr   r   r   ru   r|   r}   s   @r$   r¢   r¢   ×   sG  ø€ € € € € ðð ð ð ð ð&ð ð ð ð  ð  ð  ØØð .2Ø(,Ø37Ø26Ø04Ø26Ø!%ð\
ð \
àÔ# dÑ*ð\
ð  ™ð\
ð Ô)¨DÑ0ð	\
ð
 Ô(¨4Ñ/ð\
ð Ô&¨Ñ-ð\
ð Ô(¨4Ñ/ð\
ð ˜$‘;ð\
ð Ð+Ô,ð\
ð 
!ð\
ð \
ð \
ñ „^ñ „_ñ  Ôð\
ð \
ð \
ð \
ð \
r&   r¢   z‡
    The CTRL Model transformer with a language modeling head on top (linear layer with weights tied to the input
    embeddings).
    )Úcustom_introc                   ó$  ‡ — e Zd ZddiZˆ fd„Zee	 	 	 	 	 	 	 	 	 ddej        dz  de	dz  dej
        dz  d	ej        dz  d
ej        dz  dej
        dz  dej        dz  dedz  deej        z  dee         defd„¦   «         ¦   «         Z	 dˆ fd„	Zˆ xZS )ÚCTRLLMHeadModelzlm_head.weightztransformer.w.weightc                 óæ   •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          j        |j        |j        d¬¦  «        | _        |  	                    ¦   «          d S )NT©Úbias)
rR   rS   r¢   rœ   r   r\   rW   r·   Úlm_headrÀ   rÁ   s     €r$   rS   zCTRLLMHeadModel.__init__]  s`   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý$ VÑ,Ô,ˆÔÝ”y ¤°Ô0AÈÐMÑMÔMˆŒð 	�ŠÑÔÐÐÐr&   Nr   rÈ   rÉ   r>   rÊ   rË   rÌ   ÚlabelsrÍ   Úlogits_to_keeprA   rÎ   c
           
      óT  —  | j         |f||||||dœ|
¤Ž}|d         }t          |	t          ¦  «        rt          |	 d¦  «        n|	}|                      |dd…|dd…f         ¦  «        }d}|� | j        ||fd| j        j        i|
¤Ž}t          |||j	        |j
        |j        ¬¦  «        S )ag  
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for language modeling. Note that the labels **are shifted** inside the model, i.e. you can set
            `labels = input_ids` Indices are selected in `[-100, 0, ..., config.vocab_size]` All labels set to `-100`
            are ignored (masked), the loss is only computed for labels in `[0, ..., config.vocab_size]`

        Example:

        ```python
        >>> import torch
        >>> from transformers import AutoTokenizer, CTRLLMHeadModel

        >>> tokenizer = AutoTokenizer.from_pretrained("Salesforce/ctrl")
        >>> model = CTRLLMHeadModel.from_pretrained("Salesforce/ctrl")

        >>> # CTRL was trained with control codes as the first token
        >>> inputs = tokenizer("Wikipedia The llama is", return_tensors="pt")
        >>> assert inputs["input_ids"][0, 0].item() in tokenizer.control_codes.values()

        >>> sequence_ids = model.generate(inputs["input_ids"])
        >>> sequences = tokenizer.batch_decode(sequence_ids)
        >>> sequences
        ['Wikipedia The llama is a member of the family Bovidae. It is native to the Andes of Peru,']

        >>> outputs = model(**inputs, labels=inputs["input_ids"])
        >>> round(outputs.loss.item(), 2)
        9.21

        >>> list(outputs.logits.shape)
        [1, 5, 246534]
        ```©rÉ   r>   rÊ   rË   rÌ   rÍ   r   Nr·   )ÚlossÚlogitsrÉ   r�   rž   )rœ   r¡   rZ   Úslicerè   Úloss_functionrT   r·   r   rÉ   r�   rž   )rb   rÈ   rÉ   r>   rÊ   rË   rÌ   ré   rÍ   rê   rA   Útransformer_outputsr�   Úslice_indicesrî   rí   s                   r$   ru   zCTRLLMHeadModel.forwarde  s  € ð\ /˜dÔ.Øð	
à+Ø)Ø)Ø%Ø'Øð	
ð 	
ð ð	
ð 	
Ðð ,¨AÔ.ˆå8BÀ>ÕSVÑ8WÔ8WÐk�˜~˜o¨tÑ4Ô4Ð4Ð]kˆØ—’˜m¨A¨A¨A¨}¸a¸a¸aÐ,?Ô@ÑAÔAˆàˆØÐØ%�4Ô%ØØðð ð  œ;Ô1ðð ð	ð ˆDõ &ØØØ/Ô?Ø-Ô;Ø*Ô5ð
ñ 
ô 
ð 	
r&   Fc                 óp   •—  t          ¦   «         j        |f|||dœ|¤Ž}|                     dd ¦  «         |S )N)rÉ   rÍ   Úis_first_iterationrÊ   )rR   Úprepare_inputs_for_generationÚpop)rb   rÈ   rÉ   rÍ   rô   rA   Úmodel_inputsrc   s          €r$   rõ   z-CTRLLMHeadModel.prepare_inputs_for_generation´  s\   ø€ ð
 =•u‘w”wÔ<Øð
à+ØØ1ð	
ð 
ð
 ð
ð 
ˆð 	×ÒÐ)¨4Ñ0Ô0Ð0àÐr&   )	NNNNNNNNr   )NNF)ry   rz   r{   Ú_tied_weights_keysrS   r   r   r   rß   r
   rà   rá   rZ   ÚTensorr   r   r   ru   rõ   r|   r}   s   @r$   rä   rä   T  sz  ø€ € € € € ð +Ð,BÐCÐðð ð ð ð ð Øð .2Ø(,Ø37Ø26Ø04Ø26Ø*.Ø!%Ø-.ðK
ð K
àÔ# dÑ*ðK
ð  ™ðK
ð Ô)¨DÑ0ð	K
ð
 Ô(¨4Ñ/ðK
ð Ô&¨Ñ-ðK
ð Ô(¨4Ñ/ðK
ð Ô  4Ñ'ðK
ð ˜$‘;ðK
ð ˜eœlÑ*ðK
ð Ð+Ô,ðK
ð 
 ðK
ð K
ð K
ñ „^ñ ÔðK
ð\ SXðð ð ð ð ð ð ð ð ð r&   rä   aÎ  
    The CTRL Model transformer with a sequence classification head on top (linear layer).
    [`CTRLForSequenceClassification`] uses the last token in order to do the classification, as other causal models
    (e.g. GPT-2) do. Since it does classification on the last token, it requires to know the position of the last
    token. If a `pad_token_id` is defined in the configuration, it finds the last token that is not a padding token in
    each row. If no `pad_token_id` is defined, it simply takes the last value in each row of the batch. Since it cannot
    guess the padding tokens when `inputs_embeds` are passed instead of `input_ids`, it does the same (take the last
    value in each row of the batch).
    c                   óø   ‡ — e Zd Zˆ fd„Zee	 	 	 	 	 	 	 	 ddej        dz  dedz  dej	        dz  dej        dz  dej        dz  dej	        dz  d	ej        dz  d
e
dz  dee         defd„¦   «         ¦   «         Zˆ xZS )ÚCTRLForSequenceClassificationc                 óþ   •— t          ¦   «                              |¦  «         |j        | _        t          |¦  «        | _        t          j        |j        | j        d¬¦  «        | _        |  	                    ¦   «          d S )NFræ   )
rR   rS   Ú
num_labelsr¢   rœ   r   r\   rW   Ú
classifierrÀ   rÁ   s     €r$   rS   z&CTRLForSequenceClassification.__init__Ó  si   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø Ô+ˆŒÝ$ VÑ,Ô,ˆÔÝœ) F¤M°4´?ÈÐOÑOÔOˆŒð 	�ŠÑÔÐÐÐr&   NrÈ   rÉ   r>   rÊ   rË   rÌ   ré   rÍ   rA   rÎ   c	           
      ó¢  —  | j         |f||||||dœ|	¤Ž}
|
d         }|                      |¦  «        }|�|j        dd…         \  }}n|j        dd…         \  }}| j        j        €|dk    rt          d¦  «        ‚| j        j        €d}n¨|�}|| j        j        k                         |j        t          j	        ¦  «        }t          j
        |j        d         |j        t          j	        ¬¦  «        }||z                       d¦  «        }n)d}t                               | j        j        › d	�¦  «         |t          j
        ||j        ¬
¦  «        |f         }d}|��Z| j        j        €f| j        dk    rd| j        _        nN| j        dk    r7|j        t          j        k    s|j        t          j        k    rd| j        _        nd| j        _        | j        j        dk    rWt+          ¦   «         }| j        dk    r1 ||                     ¦   «         |                     ¦   «         ¦  «        }nŽ |||¦  «        }n�| j        j        dk    rGt/          ¦   «         } ||                     d| j        ¦  «        |                     d¦  «        ¦  «        }n*| j        j        dk    rt3          ¦   «         } |||¦  «        }t5          |||
j        |
j        ¬¦  «        S )a�  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).

        Example of single-label classification:

        ```python
        >>> import torch
        >>> from transformers import AutoTokenizer, CTRLForSequenceClassification

        >>> tokenizer = AutoTokenizer.from_pretrained("Salesforce/ctrl")
        >>> model = CTRLForSequenceClassification.from_pretrained("Salesforce/ctrl")

        >>> # CTRL was trained with control codes as the first token
        >>> inputs = tokenizer("Opinion My dog is cute", return_tensors="pt")
        >>> assert inputs["input_ids"][0, 0].item() in tokenizer.control_codes.values()

        >>> with torch.no_grad():
        ...     logits = model(**inputs).logits

        >>> predicted_class_id = logits.argmax().item()
        >>> model.config.id2label[predicted_class_id]
        'LABEL_0'
        ```

        ```python
        >>> import torch

        >>> torch.manual_seed(42)  # doctest: +IGNORE_RESULT
        >>> # To train a model on `num_labels` classes, you can pass `num_labels=num_labels` to `.from_pretrained(...)`
        >>> num_labels = len(model.config.id2label)
        >>> model = CTRLForSequenceClassification.from_pretrained("Salesforce/ctrl", num_labels=num_labels)

        >>> labels = torch.tensor(1)
        >>> loss = model(**inputs, labels=labels).loss
        >>> round(loss.item(), 2)
        0.93
        ```

        Example of multi-label classification:

        ```python
        >>> import torch
        >>> from transformers import AutoTokenizer, CTRLForSequenceClassification

        >>> tokenizer = AutoTokenizer.from_pretrained("Salesforce/ctrl")
        >>> model = CTRLForSequenceClassification.from_pretrained(
        ...     "Salesforce/ctrl", problem_type="multi_label_classification"
        ... )

        >>> # CTRL was trained with control codes as the first token
        >>> inputs = tokenizer("Opinion My dog is cute", return_tensors="pt")
        >>> assert inputs["input_ids"][0, 0].item() in tokenizer.control_codes.values()

        >>> with torch.no_grad():
        ...     logits = model(**inputs).logits

        >>> predicted_class_id = logits.argmax().item()
        >>> model.config.id2label[predicted_class_id]
        'LABEL_0'
        ```

        ```python
        >>> # To train a model on `num_labels` classes, you can pass `num_labels=num_labels` to `.from_pretrained(...)`
        >>> num_labels = len(model.config.id2label)
        >>> model = CTRLForSequenceClassification.from_pretrained("Salesforce/ctrl", num_labels=num_labels)

        >>> num_labels = len(model.config.id2label)
        >>> labels = torch.nn.functional.one_hot(torch.tensor([predicted_class_id]), num_classes=num_labels).to(
        ...     torch.float
        ... )
        >>> loss = model(**inputs, labels=labels).loss
        >>> loss.backward()  # doctest: +IGNORE_RESULT
        ```rì   r   Nr   r   z=Cannot handle batch sizes > 1 if no padding token is defined.r)   rÒ   zŠ will not detect padding tokens in `inputs_embeds`. Results may be unexpected if using padding tokens in conjunction with `inputs_embeds.`)rÐ   Ú
regressionÚsingle_label_classificationÚmulti_label_classification)rí   rî   r�   rž   )rœ   rþ   re   rT   Úpad_token_idrÔ   r.   rÐ   r   Úint32r,   ÚargmaxÚloggerÚwarning_oncerc   ry   Úproblem_typerý   r(   rÖ   rZ   r   Úsqueezer   rf   r   r   r�   rž   )rb   rÈ   rÉ   r>   rÊ   rË   rÌ   ré   rÍ   rA   rñ   r�   rî   rÚ   Úsequence_lengthÚlast_non_pad_tokenÚnon_pad_maskÚtoken_indicesÚpooled_logitsrí   Úloss_fcts                        r$   ru   z%CTRLForSequenceClassification.forwardÜ  s  € ðv /˜dÔ.Øð	
à+Ø)Ø)Ø%Ø'Øð	
ð 	
ð ð	
ð 	
Ðð ,¨AÔ.ˆØ—’ Ñ/Ô/ˆàÐ Ø*3¬/¸"¸1¸"Ô*=Ñ'ˆJ˜˜à*7Ô*=¸b¸q¸bÔ*AÑ'ˆJ˜àŒ;Ô#Ð+°
¸a²°ÝÐ\Ñ]Ô]Ð]ØŒ;Ô#Ð+Ø!#ÐÐØÐ"à%¨¬Ô)AÒA×EÒEÀfÄmÕUZÔU`ÑaÔaˆLÝ!œL¨¬¸Ô)<ÀVÄ]ÕZ_ÔZeÐfÑfÔfˆMØ"/°,Ñ">×!FÒ!FÀrÑ!JÔ!JÐÐà!#ÐÝ×ÒØ”>Ô*ð Zð Zð Zñô ð ð
 �uœ|¨J¸v¼}ÐMÑMÔMÐOaÐaÔbˆàˆØÑØŒ{Ô'Ð/Ø”? aÒ'Ð'Ø/;�D”KÔ,Ð,Ø”_ qÒ(Ð(¨f¬l½e¼jÒ.HÐ.HÈFÌLÕ\aÔ\eÒLeÐLeØ/L�D”KÔ,Ð,à/K�D”KÔ,àŒ{Ô'¨<Ò7Ð7Ý"™9œ9�Ø”? aÒ'Ð'Ø#˜8 M×$9Ò$9Ñ$;Ô$;¸V¿^º^Ñ=MÔ=MÑNÔN�D�Dà#˜8 M°6Ñ:Ô:�D�DØ”Ô)Ð-JÒJÐJÝ+Ñ-Ô-�Ø�x × 2Ò 2°2°t´Ñ GÔ GÈÏÊÐUWÉÌÑYÔY��Ø”Ô)Ð-IÒIÐIÝ,Ñ.Ô.�Ø�x ¨vÑ6Ô6�Ý'ØØ Ø-Ô;Ø*Ô5ð	
ñ 
ô 
ð 	
r&   )NNNNNNNN)ry   rz   r{   rS   r   r   r   rß   r
   rà   rá   r   r   r   ru   r|   r}   s   @r$   rû   rû   Ç  s5  ø€ € € € € ðð ð ð ð ð Øð .2Ø(,Ø37Ø26Ø04Ø26Ø*.Ø!%ðY
ð Y
àÔ# dÑ*ðY
ð  ™ðY
ð Ô)¨DÑ0ð	Y
ð
 Ô(¨4Ñ/ðY
ð Ô&¨Ñ-ðY
ð Ô(¨4Ñ/ðY
ð Ô  4Ñ'ðY
ð ˜$‘;ðY
ð Ð+Ô,ðY
ð 
"ðY
ð Y
ð Y
ñ „^ñ ÔðY
ð Y
ð Y
ð Y
ð Y
r&   rû   )rû   rä   r¢   r›   )Nr9   ):Ú__doc__Úcollections.abcr   Únumpyr×   r   r   Útorch.nnr   r   r   Ú r	   r£   Úcache_utilsr
   r   Ú
generationr   Úmasking_utilsr   Úmodeling_outputsr   r   r   Úmodeling_utilsr   r   Úprocessing_utilsr   Úutilsr   r   r   r   Úutils.genericr   Úutils.output_capturingr   Úconfiguration_ctrlr   Ú
get_loggerry   r  r%   r8   ÚModulerù   r¦   rN   rP   r‚   r„   r›   r¢   rä   rû   Ú__all__© r&   r$   ú<module>r#     sÎ  ðð Ð à $Ð $Ð $Ð $Ð $Ð $à Ð Ð Ð Ø €€€Ø Ð Ð Ð Ð Ð Ø AÐ AÐ AÐ AÐ AÐ AÐ AÐ AÐ AÐ Aà &Ð &Ð &Ð &Ð &Ð &Ø .Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ø )Ð )Ð )Ð )Ð )Ð )Ø /Ð /Ð /Ð /Ð /Ð /Ø iÐ iÐ iÐ iÐ iÐ iÐ iÐ iÐ iÐ iØ FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &ðð ð ð ð ð ð ð ð ð ð ð ð 8Ð 7Ð 7Ð 7Ð 7Ð 7Ø 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø *Ð *Ð *Ð *Ð *Ð *ð 
ˆÔ	˜HÑ	%Ô	%€ðð ð ð
ð ð ð, !Øð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð �T‰\ð%ð ð%ð Ð'Ô(ð%ð %ð %ð %ð85)ð 5)ð 5)ð 5)ð 5)˜œñ 5)ô 5)ð 5)ðp`ð `ð `ð%ð %ð %ð %ð %�2”9ñ %ô %ð %ðP ðð ð ð ð ˜/ñ ô ñ „ðð( ðy
ð y
ð y
ð y
ð y
Ð#ñ y
ô y
ñ „ðy
ðx €ððñ ô ðjð jð jð jð jÐ)¨?ñ jô jñô ðjðZ €ðð
ñ 
ô 
ðe
ð e
ð e
ð e
ð e
Ð$7ñ e
ô e
ñ
ô 
ðe
ðP cÐ
bÐ
b€€€r&   