§
    ‚Štjª  ã                   óè   — d dl mZ d dlZddl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 d
dlmZ  ej        e¦  «        Ze G d„ de¦  «        ¦   «         Ze
 G d„ de¦  «        ¦   «         ZdgZdS )é    )Ú	dataclassNé   )ÚCache)Ú$ImageClassifierOutputWithNoAttention)ÚPreTrainedModel)Úauto_docstringÚcan_return_tupleÚloggingé   )ÚAutoModelForImageTextToTexté   )ÚShieldGemma2Configc                   ó2   — e Zd ZU dZdZej        dz  ed<   dS )Ú0ShieldGemma2ImageClassifierOutputWithNoAttentionz^ShieldGemma2 classifies imags as violative or not relative to a specific policy
    Args:
    NÚprobabilities)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚtorchÚTensorÚ__annotations__© ó    út/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/shieldgemma2/modeling_shieldgemma2.pyr   r   "   s5   € € € € € € ðð ð *.€M�5”< $Ñ&Ð-Ð-Ñ-Ð-Ð-r   r   c                   óz  ‡ — e Zd ZU eed<   dZdZdZdZdZ	dZ
defˆ fd„Zd„ Zd„ Zd„ Zd	„ Zee	 	 	 	 	 	 	 	 	 	 	 	 	 d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j        d
z  dej        d
z  dej        d
z  ded
z  ded
z  ded
z  ded
z  deej        z  defd„¦   «         ¦   «         Zˆ xZS )Ú"ShieldGemma2ForImageClassificationÚconfig)ÚimageÚtextÚmodelTc                 ó   •— t          ¦   «                              |¬¦  «         t          |dd¦  «        | _        t          |dd¦  «        | _        t          j        |¬¦  «        | _        |                      ¦   «          d S )N)r   Úyes_token_indexi *  Úno_token_indexi»  )	ÚsuperÚ__init__Úgetattrr#   r$   r   Úfrom_configr!   Ú	post_init)Úselfr   Ú	__class__s     €r   r&   z+ShieldGemma2ForImageClassification.__init__5   ss   ø€ Ý‰Œ×Ò ÐÑ'Ô'Ð'Ý& vÐ/@À&ÑIÔIˆÔÝ% fÐ.>ÀÑEÔEˆÔÝ0Ô<ÀFÐKÑKÔKˆŒ
Ø�ŠÑÔÐÐÐr   c                 óX   — | j                              ¦   «                              ¦   «         S ©N)r!   Úget_decoderÚget_input_embeddings©r*   s    r   r/   z7ShieldGemma2ForImageClassification.get_input_embeddings<   s"   € ØŒz×%Ò%Ñ'Ô'×<Ò<Ñ>Ô>Ð>r   c                 ó^   — | j                              ¦   «                              |¦  «         d S r-   )r!   r.   Úset_input_embeddings)r*   Úvalues     r   r2   z7ShieldGemma2ForImageClassification.set_input_embeddings?   s*   € ØŒ
×ÒÑ Ô ×5Ò5°eÑ<Ô<Ð<Ð<Ð<r   c                 óX   — | j                              ¦   «                              ¦   «         S r-   )r!   r.   Úget_output_embeddingsr0   s    r   r5   z8ShieldGemma2ForImageClassification.get_output_embeddingsB   s"   € ØŒz×%Ò%Ñ'Ô'×=Ò=Ñ?Ô?Ð?r   c                 ó^   — | j                              ¦   «                              |¦  «         d S r-   )r!   r.   Úset_output_embeddings)r*   Únew_embeddingss     r   r7   z8ShieldGemma2ForImageClassification.set_output_embeddingsE   s*   € ØŒ
×ÒÑ Ô ×6Ò6°~ÑFÔFÐFÐFÐFr   Nr   Ú	input_idsÚpixel_valuesÚattention_maskÚposition_idsÚpast_key_valuesÚtoken_type_idsÚinputs_embedsÚlabelsÚ	use_cacheÚoutput_attentionsÚoutput_hidden_statesÚreturn_dictÚlogits_to_keepÚreturnc                 óÆ   —  | j         d|||||||||	|
|||dœ|¤Ž}|j        }|dd…d| j        | j        gf         }t	          j        |d¬¦  «        }t          ||¬¦  «        S )aY  
        Returns:
            A `ShieldGemma2ImageClassifierOutputWithNoAttention` instance containing the logits and probabilities
            associated with the model predicting the `Yes` or `No` token as the response to that prompt, captured in the
            following properties.

                *   `logits` (`torch.Tensor` of shape `(batch_size, 2)`):
                    The first position along dim=1 is the logits for the `Yes` token and the second position along dim=1 is
                    the logits for the `No` token.
                *   `probabilities` (`torch.Tensor` of shape `(batch_size, 2)`):
                    The first position along dim=1 is the probability of predicting the `Yes` token and the second position
                    along dim=1 is the probability of predicting the `No` token.

            ShieldGemma prompts are constructed such that predicting the `Yes` token means the content *does violate* the
            policy as described. If you are only interested in the violative condition, use
            `violated = outputs.probabilities[:, 1]` to extract that slice from the output tensors.

            When used with the `ShieldGemma2Processor`, the `batch_size` will be equal to `len(images) * len(policies)`,
            and the order within the batch will be img1_policy1, ... img1_policyN, ... imgM_policyN.
        )r9   r:   r;   r<   r=   r>   r?   r@   rA   rB   rC   rD   rE   Néÿÿÿÿ)Údim)Úlogitsr   r   )r!   rJ   r#   r$   r   Úsoftmaxr   )r*   r9   r:   r;   r<   r=   r>   r?   r@   rA   rB   rC   rD   rE   Ú	lm_kwargsÚoutputsrJ   Úselected_logitsr   s                      r   Úforwardz*ShieldGemma2ForImageClassification.forwardH   s­   € ðN �$”*ð 
ØØ%Ø)Ø%Ø+Ø)Ø'ØØØ/Ø!5Ø#Ø)ð
ð 
ð ð
ð 
ˆð  ”ˆØ     B¨Ô)=¸tÔ?RÐ(SÐ!SÔTˆÝœ o¸2Ð>Ñ>Ô>ˆÝ?Ø"Ø'ð
ñ 
ô 
ð 	
r   )NNNNNNNNNNNNr   )r   r   r   r   r   Úinput_modalitiesÚbase_model_prefixÚ_supports_flash_attnÚ_supports_sdpaÚ_supports_flex_attnÚ_supports_attention_backendr&   r/   r2   r5   r7   r   r	   r   Ú
LongTensorÚFloatTensorr   r   ÚboolÚintr   rO   Ú__classcell__)r+   s   @r   r   r   +   sí  ø€ € € € € € àÐÐÑØ(ÐØÐØÐØ€NØÐØ"&ÐðÐ1ð ð ð ð ð ð ð?ð ?ð ?ð=ð =ð =ð@ð @ð @ðGð Gð Gð Øð .2Ø15Ø.2Ø04Ø(,Ø26Ø26Ø*.Ø!%Ø)-Ø,0Ø#'Ø-.ð;
ð ;
àÔ# dÑ*ð;
ð Ô'¨$Ñ.ð;
ð œ tÑ+ð	;
ð
 Ô&¨Ñ-ð;
ð  ™ð;
ð Ô(¨4Ñ/ð;
ð Ô(¨4Ñ/ð;
ð Ô  4Ñ'ð;
ð ˜$‘;ð;
ð   $™;ð;
ð # T™kð;
ð ˜D‘[ð;
ð ˜eœlÑ*ð;
ð  
:ð!;
ð ;
ð ;
ñ Ôñ „^ð;
ð ;
ð ;
ð ;
ð ;
r   r   )Údataclassesr   r   Úcache_utilsr   Úmodeling_outputsr   Úmodeling_utilsr   Úutilsr   r	   r
   Úautor   Úconfiguration_shieldgemma2r   Ú
get_loggerr   Úloggerr   r   Ú__all__r   r   r   ú<module>re      sJ  ðð "Ð !Ð !Ð !Ð !Ð !à €€€à  Ð  Ð  Ð  Ð  Ð  Ø DÐ DÐ DÐ DÐ DÐ DØ -Ð -Ð -Ð -Ð -Ð -ðð ð ð ð ð ð ð ð ð ð
 /Ð .Ð .Ð .Ð .Ð .Ø :Ð :Ð :Ð :Ð :Ð :ð 
ˆÔ	˜HÑ	%Ô	%€ð ð.ð .ð .ð .ð .Ð7[ñ .ô .ñ „ð.ð ðY
ð Y
ð Y
ð Y
ð Y
¨ñ Y
ô Y
ñ „ðY
ðz )ð€€€r   