
    ^j                     2   d dl mZ d dlmZ 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 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 ddlmZmZm Z m!Z!m"Z" ddl#m$Z$ ddl%m&Z& ddl'm(Z(m)Z)m*Z* e e G d de                    Z+ G d dejX                        Z-	 dNdejX                  dej\                  dej\                  dej\                  dej\                  dz  de/de/fdZ0 G d d ejX                        Z1 G d! d"ejX                        Z2 G d# d$ejX                        Z3 G d% d&ejX                        Z4 G d' d(ejX                        Z5 G d) d*e      Z6 G d+ d,ejX                        Z7 G d- d.ejX                        Z8 G d/ d0ejX                        Z9 G d1 d2ejX                        Z: G d3 d4e      Z; G d5 d6ejX                        Z< G d7 d8ejX                        Z=e  G d9 d:e             Z> e d;<       G d= d>e>             Z? e d?<       G d@ dAe>             Z@ G dB dCe>      ZAdDej\                  dEej\                  fdFZBdGej\                  dEej\                  fdHZCdIej\                  dEej\                  fdJZDe  G dK dLe>             ZEg dMZFy)O    )Callable)	dataclass)AnyN   )initialization)ACT2FN)create_bidirectional_mask)GradientCheckpointingLayer)BaseModelOutputBaseModelOutputWithPooling'BaseModelOutputWithPoolingAndProjection)ALL_ATTENTION_FUNCTIONSPreTrainedModel)Unpack)apply_chunking_to_forward)ModelOutputTransformersKwargsauto_docstringcan_return_tuple	torch_int)merge_with_config_defaults)capture_outputs   )AltCLIPConfigAltCLIPTextConfigAltCLIPVisionConfigc                      e Zd ZU dZdZej                  dz  ed<   dZej                  dz  ed<   dZ	ej                  dz  ed<   dZ
ej                  dz  ed<   dZej                  dz  ed<   dZeed<   dZeed	<   d
ee   fdZy)AltCLIPOutputa  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `return_loss` is `True`):
        Contrastive loss for image-text similarity.
    logits_per_image (`torch.FloatTensor` of shape `(image_batch_size, text_batch_size)`):
        The scaled dot product scores between `image_embeds` and `text_embeds`. This represents the image-text
        similarity scores.
    logits_per_text (`torch.FloatTensor` of shape `(text_batch_size, image_batch_size)`):
        The scaled dot product scores between `text_embeds` and `image_embeds`. This represents the text-image
        similarity scores.
    text_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim`):
        The text embeddings obtained by applying the projection layer to the pooled output of [`AltCLIPTextModel`].
    image_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim`):
        The image embeddings obtained by applying the projection layer to the pooled output of [`AltCLIPVisionModel`].
    text_model_output (`BaseModelOutputWithPooling`):
        The output of the [`AltCLIPTextModel`].
    vision_model_output (`BaseModelOutputWithPooling`):
        The output of the [`AltCLIPVisionModel`].
    Nlosslogits_per_imagelogits_per_texttext_embedsimage_embedstext_model_outputvision_model_outputreturnc                 B    t        d | j                         D              S )Nc              3   `   K   | ]&  }t        |t              r|j                         n| ( y wN)
isinstancer   to_tuple).0vs     w/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/altclip/modeling_altclip.py	<genexpr>z)AltCLIPOutput.to_tuple.<locals>.<genexpr>H   s$     ^1Z;%?QZZ\QF^s   ,.)tuplevalues)selfs    r.   r+   zAltCLIPOutput.to_tupleG   s    ^PTP[P[P]^^^    )__name__
__module____qualname____doc__r   torchFloatTensor__annotations__r    r!   r"   r#   r$   r   r%   r0   r   r+    r3   r.   r   r   )   s    & &*D%

d
")15e''$.504OU&&-4,0K""T)0-1L%##d*148186:3:_%* _r3   r   c                        e Zd ZdZ fdZ	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  ded	ej                  fd
Z
ed        Zedd       Z xZS )AltRobertaEmbeddingszGConstruct the embeddings from word, position and token_type embeddings.c                 T   t         |           t        j                  |j                  |j
                  |j                        | _        t        j                  |j                  |j
                        | _	        t        j                  |j
                  |j                        | _
        t        j                  |j                        | _        | j                  dt!        j"                  |j$                        j'                  d      d       | j                  dt!        j(                  | j*                  j-                         t         j.                        d       |j                  | _        t        j                  |j$                  |j
                  | j0                        | _        y )	N)padding_idxepsposition_idsr   F
persistenttoken_type_idsdtype)super__init__nn	Embedding
vocab_sizehidden_sizepad_token_idword_embeddingstype_vocab_sizetoken_type_embeddings	LayerNormlayer_norm_epsDropouthidden_dropout_probdropoutregister_bufferr8   arangemax_position_embeddingsexpandzerosrB   sizelongr?   position_embeddingsr2   config	__class__s     r.   rK   zAltRobertaEmbeddings.__init__N   s4   !||F,=,=v?Q?Q_e_r_rs%'\\&2H2H&J\J\%]"f&8&8f>S>STzz&"<"<=ELL)G)GHOOPWXej 	 	
 	ekk$*;*;*@*@*B%**Ubg 	 	
 "..#%<<**F,>,>DL\L\$
 r3   N	input_idsrG   rB   inputs_embedspast_key_values_lengthr&   c                    |<|| j                  || j                  |      }n| j                  || j                        }||j                         }n|j                         d d }|\  }}|t	        | d      rm| j
                  j                  |j                        j                  |j                  d   d      }	t        j                  |	d|      }	|	j                  ||      }n:t        j                  |t        j                  | j                  j                        }|| j                  |      }| j!                  |      }
||
z   }| j#                  |      }||z   }| j%                  |      }| j'                  |      }|S )NrD   rG   r   r   )dimindexrI   device)"create_position_ids_from_input_idsr?   &create_position_ids_from_inputs_embedsr^   hasattrrG   tork   r\   shaper8   gatherr]   r_   rB   rQ   rS   r`   rT   rX   )r2   rd   rG   rB   re   rf   input_shape
batch_size
seq_lengthbuffered_token_type_idsrS   
embeddingsr`   s                r.   forwardzAltRobertaEmbeddings.forwardb   s    $#FFt//1G   $JJ=Z^ZjZjk #..*K',,.s3K!,
J
 !t-.*.*=*=*@*@ATAT*U*\*\]i]o]opq]rtv*w'*/,,7NTU]i*j'!8!?!?
J!W!&[

SWSdSdSkSk!l  00;M $ : :> J"%::
"66|D"55
^^J/
\\*-
r3   c                     | j                         dd }|d   }t        j                  |dz   ||z   dz   t        j                  | j                        }|j                  d      j                  |      S )z
        We are provided embeddings directly. We cannot infer which are padded so just generate sequential position ids.

        Args:
            inputs_embeds: torch.Tensor

        Returns: torch.Tensor
        NrD   r   rj   r   )r^   r8   rZ   r_   rk   	unsqueezer\   )re   r?   rr   sequence_lengthrB   s        r.   rm   z;AltRobertaEmbeddings.create_position_ids_from_inputs_embeds   sp     $((*3B/%a.||!O_{:Q>ejjYfYmYm
 %%a(//<<r3   c                     | j                  |      j                         }t        j                  |d      j	                  |      |z   |z  }|j                         |z   S )a  
        Replace non-padding symbols with their position numbers. Position numbers begin at padding_idx+1. Padding symbols
        are ignored. This is modified from fairseq's `utils.make_positions`.

        Args:
            x: torch.Tensor x:

        Returns: torch.Tensor
        r   rh   )neintr8   cumsumtype_asr_   )rd   r?   rf   maskincremental_indicess        r.   rl   z7AltRobertaEmbeddings.create_position_ids_from_input_ids   sW     ||K(,,.$||Da8@@FI__cgg"'')K77r3   )NNNNr   )r   )r4   r5   r6   r7   rK   r8   
LongTensorr9   r~   Tensorrw   staticmethodrm   rl   __classcell__rc   s   @r.   r=   r=   K   s    Q
, .2260426&'.##d*. ((4/. &&-	.
 ((4/. !$. 
.` = =" 8 8r3   r=   modulequerykeyvalueattention_maskscalingrX   c                    t        j                  ||j                  dd            |z  }|||z   }t        j                  j                  |dt         j                        j                  |j                        }t        j                  j                  ||| j                        }t        j                  ||      }	|	j                  dd      j                         }	|	|fS )N   r   rD   )rh   rI   )ptrainingr   )r8   matmul	transposerL   
functionalsoftmaxfloat32ro   rI   rX   r   
contiguous)
r   r   r   r   r   r   rX   kwargsattn_weightsattn_outputs
             r.   eager_attention_forwardr      s     <<s}}Q':;gEL!#n4==((2U]](SVVW\WbWbcL==((6??([L,,|U3K''1-88:K$$r3   c                        e Zd Z fdZ	 ddej
                  dej                  dz  dee   de	ej
                  ej
                  dz  f   fdZ
 xZS )	AltRobertaSelfAttentionc                 2   t         |           |j                  |j                  z  dk7  r2t	        |d      s&t        d|j                   d|j                   d      || _        |j                  | _        t        |j                  |j                  z        | _        | j                  | j                  z  | _	        t        j                  |j                  | j                        | _        t        j                  |j                  | j                        | _        t        j                  |j                  | j                        | _        t        j                  |j                         | _        |j                   | _        | j                  dz  | _        d| _        y )Nr   embedding_sizezThe hidden size (z6) is not a multiple of the number of attention heads ()      F)rJ   rK   rO   num_attention_headsrn   
ValueErrorrb   r~   attention_head_sizeall_head_sizerL   Linearr   r   r   rV   attention_probs_dropout_probrX   attention_dropoutr   	is_causalra   s     r.   rK   z AltRobertaSelfAttention.__init__   sJ    : ::a?PVXhHi#F$6$6#7 8 445Q8 
 #)#=#= #&v'9'9F<V<V'V#W !558P8PPYYv1143E3EF
99V//1C1CDYYv1143E3EF
zz&"E"EF!'!D!D//5r3   Nhidden_statesr   r   r&   c                 x   |j                   d d }g |d| j                  }| j                  |      j                  |      j	                  dd      }| j                  |      j                  |      j	                  dd      }| j                  |      j                  |      j	                  dd      }t        j                  | j                  j                  t              }	 |	| ||||f| j                  sdn| j                  | j                  d|\  }
} |
j                  g |d j!                         }
|
|fS )NrD   r   r           )rX   r   )rp   r   r   viewr   r   r   r   get_interfacerb   _attn_implementationr   r   r   r   reshaper   )r2   r   r   r   rr   hidden_shapequery_states
key_statesvalue_statesattention_interfacer   r   s               r.   rw   zAltRobertaSelfAttention.forward   s>    $))#2.CCbC$*B*BCzz-055lCMMaQRSXXm,11,?II!QO
zz-055lCMMaQRS(?(M(MKK,,.E)
 %8	%
  $}}C$2H2HLL	%
 	%
!\ *k));;;;FFHL((r3   r)   )r4   r5   r6   rK   r8   r   r9   r   r   r0   rw   r   r   s   @r.   r   r      sd    2 48)||) ))D0) +,	)
 
u||U\\D00	1)r3   r   c                   n     e Zd Z fdZdej
                  dej
                  dej
                  fdZ xZS )AltRobertaSelfOutputc                 (   t         |           t        j                  |j                  |j                        | _        t        j                  |j                  |j                        | _        t        j                  |j                        | _
        y Nr@   )rJ   rK   rL   r   rO   denserT   rU   rV   rW   rX   ra   s     r.   rK   zAltRobertaSelfOutput.__init__  s`    YYv1163E3EF
f&8&8f>S>STzz&"<"<=r3   r   input_tensorr&   c                 r    | j                  |      }| j                  |      }| j                  ||z         }|S r)   r   rX   rT   r2   r   r   s      r.   rw   zAltRobertaSelfOutput.forward	  7    

=1]3}|'CDr3   r4   r5   r6   rK   r8   r   rw   r   r   s   @r.   r   r     1    >U\\  RWR^R^ r3   r   c            	            e Zd Z fdZ	 ddej
                  dej                  dz  dee   dej
                  fdZ	 xZ
S )	AltRobertaAttentionc                 b    t         |           t        |      | _        t	        |      | _        y r)   )rJ   rK   r   r2   r   outputra   s     r.   rK   zAltRobertaAttention.__init__  s&    +F3	*62r3   Nr   r   r   r&   c                 ^    |} | j                   |fd|i|\  }}| j                  ||      }|S Nr   )r2   r   r2   r   r   r   residual_s         r.   rw   zAltRobertaAttention.forward  sK     !$499
)
 
q
 M8<r3   r)   )r4   r5   r6   rK   r8   r   r9   r   r   rw   r   r   s   @r.   r   r     sQ    3 48|| ))D0 +,	
 
r3   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )AltRobertaIntermediatec                    t         |           t        j                  |j                  |j
                        | _        t        |j                  t              rt        |j                     | _        y |j                  | _        y r)   )rJ   rK   rL   r   rO   intermediate_sizer   r*   
hidden_actstrr   intermediate_act_fnra   s     r.   rK   zAltRobertaIntermediate.__init__'  s]    YYv1163K3KL
f''-'-f.?.?'@D$'-'8'8D$r3   r   r&   c                 J    | j                  |      }| j                  |      }|S r)   )r   r   r2   r   s     r.   rw   zAltRobertaIntermediate.forward/  s&    

=100?r3   r   r   s   @r.   r   r   &  s#    9U\\ ell r3   r   c                   n     e Zd Z fdZdej
                  dej
                  dej
                  fdZ xZS )AltRobertaOutputc                 (   t         |           t        j                  |j                  |j
                        | _        t        j                  |j
                  |j                        | _        t        j                  |j                        | _        y r   )rJ   rK   rL   r   r   rO   r   rT   rU   rV   rW   rX   ra   s     r.   rK   zAltRobertaOutput.__init__6  s`    YYv779K9KL
f&8&8f>S>STzz&"<"<=r3   r   r   r&   c                 r    | j                  |      }| j                  |      }| j                  ||z         }|S r)   r   r   s      r.   rw   zAltRobertaOutput.forward<  r   r3   r   r   s   @r.   r   r   5  r   r3   r   c            	            e Zd Z fdZ	 d	dej
                  dej                  dz  dee   dej
                  fdZ	d Z
 xZS )
AltRobertaLayerc                     t         |           |j                  | _        d| _        t	        |      | _        t        |      | _        t        |      | _	        y )Nr   )
rJ   rK   chunk_size_feed_forwardseq_len_dimr   	attentionr   intermediater   r   ra   s     r.   rK   zAltRobertaLayer.__init__D  sI    '-'E'E$,V426:&v.r3   Nr   r   r   r&   c                      | j                   |fd|i|}t        | j                  | j                  | j                  |      }|S r   )r   r   feed_forward_chunkr   r   )r2   r   r   r   s       r.   rw   zAltRobertaLayer.forwardL  sY     '
)
 
 2##T%A%A4CSCSUb
 r3   c                 L    | j                  |      }| j                  ||      }|S r)   )r   r   )r2   attention_outputintermediate_outputlayer_outputs       r.   r   z"AltRobertaLayer.feed_forward_chunk^  s,    "//0@A{{#68HIr3   r)   )r4   r5   r6   rK   r8   r   r9   r   r   rw   r   r   r   s   @r.   r   r   C  sV    / 48|| ))D0 +,	
 
$r3   r   c                   `     e Zd ZdZdef fdZ	 d	dej                  dz  dee	   de
fdZ xZS )
AltRobertaEncoderz
    Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a
    [`AltRobertaEncoderLayer`].

    Args:
        config: AltCLIPTextConfig
    rb   c                     t         |           || _        t        j                  t        |j                        D cg c]  }t        |       c}      | _        d| _	        y c c}w NF)
rJ   rK   rb   rL   
ModuleListrangenum_hidden_layersr   layersgradient_checkpointingr2   rb   r   rc   s      r.   rK   zAltRobertaEncoder.__init__m  sN    mmeFLdLdFe$f_V%<$fg&+# %g   A#Nr   r   r&   c                 T    |}| j                   D ]  } |||fi |} t        |      S N)last_hidden_stater   r   r2   re   r   r   r   encoder_layers         r.   rw   zAltRobertaEncoder.forwards  H     &![[ 	M) M	 +
 	
r3   r)   )r4   r5   r6   r7   r   rK   r8   r   r   r   r   rw   r   r   s   @r.   r   r   d  sL    ,0 , /3
 t+
 +,	

 

r3   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )AltRobertaPoolerc                     t         |           t        j                  |j                  |j                        | _        t        j                         | _        y r)   )rJ   rK   rL   r   rO   r   Tanh
activationra   s     r.   rK   zAltRobertaPooler.__init__  s9    YYv1163E3EF
'')r3   r   r&   c                 \    |d d df   }| j                  |      }| j                  |      }|S )Nr   )r   r   )r2   r   first_token_tensorpooled_outputs       r.   rw   zAltRobertaPooler.forward  s6     +1a40

#566r3   r   r   s   @r.   r   r     s#    $
U\\ ell r3   r   c                        e Zd ZdZdeez  f fdZ	 d
dej                  dej                  dz  de	e
   deej                  ej                  dz  f   fd	Z xZS )AltCLIPAttentionz=Multi-headed attention from 'Attention Is All You Need' paperrb   c                    t         |           || _        |j                  | _        |j
                  | _        | j                  | j                  z  | _        | j                  dz  | _        |j                  | _
        d| _        t        j                  | j                  | j                        | _        t        j                  | j                  | j                        | _        t        j                  | j                  | j                        | _        t        j                  | j                  | j                        | _        y )Nr   F)rJ   rK   rb   rO   	embed_dimr   	num_headshead_dimscaler   rX   r   rL   r   k_projv_projq_projout_projra   s     r.   rK   zAltCLIPAttention.__init__  s    ++33$..8]]D(
//ii?ii?ii?		$..$..Ar3   Nr   r   r   r&   c                    |j                   dd }g |d| j                  }| j                  |      }| j                  |      }| j	                  |      }|j                  |      j                  dd      }|j                  |      j                  dd      }|j                  |      j                  dd      }t        j                  | j                  j                  t              }	 |	| ||||f| j                  | j                  sdn| j                  d|\  }
} |
j                  g |d j!                         }
| j#                  |
      }
|
|fS )z#Input shape: Batch x Time x ChannelNrD   r   r   r   )r   rX   )rp   r  r  r  r  r   r   r   r   rb   r   r   r  r   rX   r   r   r	  )r2   r   r   r   rr   r   querieskeysr1   r   r   r   s               r.   rw   zAltCLIPAttention.forward  sO    $))#2.88b8$--8++m,{{=)]+,,|,66q!<yy&00A6\*44Q:(?(M(MKK,,.E)
 %8	%
 JJ#}}C$,,	%
 	%
!\ *k));;;;FFHmmK0L((r3   r)   )r4   r5   r6   r7   r   r   rK   r8   r   r   r   r0   rw   r   r   s   @r.   r   r     su    GB25FF B$ /3%)||%) t+%) +,	%)
 
u||U\\D00	1%)r3   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )
AltCLIPMLPc                    t         |           || _        t        |j                     | _        t        j                  |j                  |j                        | _
        t        j                  |j                  |j                        | _        y r)   )rJ   rK   rb   r   r   activation_fnrL   r   rO   r   fc1fc2ra   s     r.   rK   zAltCLIPMLP.__init__  sd    #F$5$5699V//1I1IJ99V55v7I7IJr3   r   r&   c                 l    | j                  |      }| j                  |      }| j                  |      }|S r)   )r  r  r  r   s     r.   rw   zAltCLIPMLP.forward  s4    /**=9/r3   r   r   s   @r.   r  r    s$    KU\\ ell r3   r  c                   ~     e Zd Zdef fdZdej                  dej                  dee   dej                  fdZ
 xZS )AltCLIPEncoderLayerrb   c                 D   t         |           |j                  | _        t	        |      | _        t        j                  | j                  |j                        | _	        t        |      | _        t        j                  | j                  |j                        | _        y r   )rJ   rK   rO   r  r   	self_attnrL   rT   rU   layer_norm1r  mlplayer_norm2ra   s     r.   rK   zAltCLIPEncoderLayer.__init__  sm    ++)&1<<F<Q<QRf%<<F<Q<QRr3   r   r   r   r&   c                     |}| j                  |      } | j                  d||d|\  }}||z   }|}| j                  |      }| j                  |      }||z   }|S )N)r   r   r;   )r  r  r  r  r   s         r.   rw   zAltCLIPEncoderLayer.forward  s     !((7)4>> 
')
 
q
 !=0 ((7/ =0r3   )r4   r5   r6   r   rK   r8   r   r   r   r9   rw   r   r   s   @r.   r  r    sQ    S2 S||  +,	
 
		r3   r  c                   `     e Zd ZdZdef fdZ	 d	dej                  dz  dee	   de
fdZ xZS )
AltCLIPEncoderz
    Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a
    [`AltCLIPEncoderLayer`].

    Args:
        config: AltCLIPConfig
    rb   c                     t         |           || _        t        j                  t        |j                        D cg c]  }t        |       c}      | _        d| _	        y c c}w r   )
rJ   rK   rb   rL   r   r   r   r  r   r   r   s      r.   rK   zAltCLIPEncoder.__init__  sP    mm%PVPhPhJi$jQ%8%@$jk&+# %kr   Nr   r   r&   c                 T    |}| j                   D ]  } |||fi |} t        |      S r   r   r   s         r.   rw   zAltCLIPEncoder.forward  r   r3   r)   )r4   r5   r6   r7   r   rK   r8   r   r   r   r   rw   r   r   s   @r.   r  r    sK    ,} , /3
 t+
 +,	

 

r3   r  c                        e Zd Zdef fdZdej                  dededej                  fdZd
dej                  dej                  fd	Z
 xZS )AltCLIPVisionEmbeddingsrb   c                    t         |           || _        |j                  | _        |j
                  | _        |j                  | _        t        j                  t        j                  | j                              | _        t        j                  |j                  | j                  | j                  | j                  d      | _        | j
                  | j                  z  dz  | _        | j                  dz   | _        t        j"                  | j                   | j                        | _        | j'                  dt        j(                  | j                         j+                  d      d       y )NF)in_channelsout_channelskernel_sizestridebiasr   r   rB   rC   rE   )rJ   rK   rb   rO   r  
image_size
patch_sizerL   	Parameterr8   randnclass_embeddingConv2dnum_channelspatch_embeddingnum_patchesnum_positionsrM   position_embeddingrY   rZ   r\   ra   s     r.   rK   z AltCLIPVisionEmbeddings.__init__"  s	   ++ ++ ++!||EKK,GH!yy++?? 
 !OOt>1D!--1"$,,t/A/A4>>"R^U\\$:L:L-M-T-TU\-]jopr3   rv   heightwidthr&   c                    |j                   d   dz
  }| j                  j                  j                  d      }|j                   d   dz
  }t        j
                  j                         s%||k(  r ||k(  r| j                  | j                        S |ddddf   }|ddddf   }|j                   d   }	|| j                  z  }
|| j                  z  }t        |dz        }|j                  d|||	      }|j                  dddd      }t        j                  j                  ||
|fdd	
      }|j                  dddd      j                  dd|	      }t	        j                   ||fd      S )a   
        This method allows to interpolate the pre-trained position encodings, to be able to use the model on higher resolution
        images. This method is also adapted to support torch.jit tracing.

        Adapted from:
        - https://github.com/facebookresearch/dino/blob/de9ee3df6cf39fac952ab558447af1fa1365362a/vision_transformer.py#L174-L194, and
        - https://github.com/facebookresearch/dinov2/blob/e1277af2ba9496fbadf7aec6eba56e8d882d1e35/dinov2/models/vision_transformer.py#L179-L211
        r   r   NrD         ?r   r   bicubicF)r^   modealign_cornersr|   )rp   r2  weightry   r8   jit
is_tracingrB   r)  r   r   permuterL   r   interpolater   cat)r2   rv   r3  r4  r0  r2  r1  class_pos_embedpatch_pos_embedrh   
new_height	new_widthsqrt_num_positionss                r.   interpolate_pos_encodingz0AltCLIPVisionEmbeddings.interpolate_pos_encoding8  sv    !&&q)A-!44;;EEaH*003a7 yy##%+*F6UZ?**4+<+<==,QU3,QU3r"t.
T__,	&}c'9:)11!5GI[]`a)11!Q1=--33i(	 4 
 *11!Q1=BB1b#Nyy/?;CCr3   pixel_valuesc                 `   |j                   \  }}}}|sJ|| j                  k7  s|| j                  k7  r,t        d| d| d| j                   d| j                   d	      | j                  j                  j
                  }| j                  |j                  |            }|j                  d      j                  dd      }| j                  j                  |dd      }	t        j                  |	|gd	      }
|r|
| j                  |
||      z   }
|
S |
| j                  | j                        z   }
|
S )
NzInput image size (*z) doesn't match model (z).rH   r   r   rD   r|   )rp   r(  r   r/  r:  rI   ro   flattenr   r,  r\   r8   r?  rE  r2  rB   )r2   rF  rE  rs   r   r3  r4  target_dtypepatch_embedsclass_embedsrv   s              r.   rw   zAltCLIPVisionEmbeddings.forwarda  s6   '3'9'9$
Avu'Vt-F%SWSbSbJb$VHAeW4KDOOK\\]^b^m^m]nnpq  ++2288++LOO,O,OP#++A.88A>++22:q"EYYl;C
##d&C&CJPVX]&^^J  $d&=&=d>O>O&PPJr3   F)r4   r5   r6   r   rK   r8   r   r~   rE  r9   rw   r   r   s   @r.   r!  r!  !  se    q2 q,'D5<< 'D 'DUX 'D]b]i]i 'DRE$5$5 Z_ZfZf r3   r!  c                        e Zd ZU eed<   dZdZg dZdZdZ	dZ
dZdZeedZ ej"                          fd       Z xZS )AltCLIPPreTrainedModelrb   altclip)imagetext)r=   r   r  r!  Tr   
attentionsc                    t         |   |       | j                  j                  }t	        |t
              rt        j                  |j                  d|j                  dz  |z         t        j                  |j                  j                  |j                  j                  |z         t        j                  |j                  j                  |j                  j                  |z         t        j                  |j                  t!        j"                  |j$                        j'                  d             yt	        |t(              r|j                  dz  d|j                  j*                  z  dz  z  |z  }|j                  dz  |z  }t        j                  |j,                  j                  |       t        j                  |j.                  j                  |       t        j                  |j0                  j                  |       t        j                  |j2                  j                  |       yt	        |t4              r|j                  j6                  dz  d|j                  j*                  z  dz  z  |z  }d|j                  j6                  z  dz  |z  }t        j                  |j8                  j                  |       t        j                  |j:                  j                  |       yt	        |t<              rwt        j                  |j>                  j                  |j@                  dz  |z         t        j                  |jB                  j                  |jD                  dz  |z         yt	        |tF              ryt        j                  |j                  t!        j"                  |j                  jH                  d         j'                  d             t        jJ                  |jL                         yy)	zInitialize the weightsr   r   )meanstd)rW  rC   r   rD   N)'rJ   _init_weightsrb   initializer_factorr*   r!  initnormal_r,  r  r/  r:  initializer_ranger2  copy_rB   r8   rZ   r1  r\   r   r   r  r  r  r	  r  rO   r  r  AltCLIPModeltext_projectiontext_embed_dimvisual_projectionvision_embed_dimr=   rp   zeros_rG   )r2   r   factorin_proj_stdout_proj_stdfc_stdrc   s         r.   rX  z$AltCLIPPreTrainedModel._init_weights  s    	f%//f56LL//cv?O?OQU?UX^?^_LL//66FMM<[<[^d<deLL2299v}}?^?^ag?ghJJv**ELL9M9M,N,U,UV],^_ 01!++T1q6==;Z;Z7Z_c6cdgmmK",,d2f<LLL--;?LL--;?LL--;?LL//\B
+!==44d:FMMDcDc@chl?lmpvvK&--333<vEFLL**7LL**<-LL&&--))4/&8 LL((//++T1F:  45JJv**ELL9L9L9R9RSU9V,W,^,^_f,ghKK--. 6r3   )r4   r5   r6   r   r:   base_model_prefixinput_modalities_no_split_modulessupports_gradient_checkpointing_supports_sdpa_supports_flash_attn_supports_flex_attn_supports_attention_backendr  r   _can_record_outputsr8   no_gradrX  r   r   s   @r.   rO  rO  t  sb    !(u&*#N"&,&
 U]]_ /  /r3   rO  zN
    The vision model from ALTCLIP without any head or projection on top.
    )custom_introc                        e Zd ZU eed<   dZdZdZdef fdZe	 e
d      e	 	 ddej                  dz  d	edz  d
ee   defd                     Z xZS )AltCLIPVisionModelrb   rF  )rQ  r/  c                 4   t         |   |       |j                  }t        |      | _        t        j                  ||j                        | _        t        |      | _
        t        j                  ||j                        | _        | j                          y r   )rJ   rK   rO   r!  rv   rL   rT   rU   pre_layrnormr  encoderpost_layernorm	post_init)r2   rb   r  rc   s      r.   rK   zAltCLIPVisionModel.__init__  so     &&	1&9LL8M8MN%f- ll9&:O:OPr3   F)tie_last_hidden_statesNrE  r   r&   c                     | j                  ||      }| j                  |      } | j                  dd|i|}|j                  }|dddddf   }| j	                  |      }t        ||      S )a  
        Examples:

        ```python
        >>> import httpx
        >>> from io import BytesIO
        >>> from PIL import Image
        >>> from transformers import AutoProcessor, AltCLIPVisionModel

        >>> model = AltCLIPVisionModel.from_pretrained("BAAI/AltCLIP")
        >>> processor = AutoProcessor.from_pretrained("BAAI/AltCLIP")

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read()))

        >>> inputs = processor(images=image, return_tensors="pt")

        >>> outputs = model(**inputs)
        >>> last_hidden_state = outputs.last_hidden_state
        >>> pooled_output = outputs.pooler_output  # pooled CLS states
        ```)rE  re   Nr   r   pooler_outputr;   )rv   rv  rw  r   rx  r   )r2   rF  rE  r   r   encoder_outputsr   r   s           r.   rw   zAltCLIPVisionModel.forward  s    > Ogh))-8+74<< ,
',
,

 ,==)!Q'2++M:)/'
 	
r3   r   )r4   r5   r6   r   r:   main_input_nameri  _input_embed_layerrK   r   r   r   r8   r9   boolr   r   r   rw   r   r   s   @r.   rt  rt    s      $O!*2   E2 2605+
''$.+
 #'++
 +,	+

 
$+
  3  +
r3   rt  aE  
    The model behaves as an encoder following the architecture described in *Attention is
    all you need*_ by Ashish Vaswani, Noam Shazeer, Niki Parmar, Jakob Uszkoreit, Llion Jones, Aidan N. Gomez, Lukasz
    Kaiser and Illia Polosukhin.

    .. _*Attention is all you need*: https://huggingface.co/papers/1706.03762
    c                       e Zd ZU eed<   dZdZeedZ	d fd	Z
e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j                  dz  dee   deez  fd                     Z xZS )AltRobertaModelrb   rR  rQ   rS  c                     t         |   |       t        |      | _        t	        |      | _        |rt        |      nd| _        | j                          y)zv
        add_pooling_layer (bool, *optional*, defaults to `True`):
            Whether to add a pooling layer
        N)	rJ   rK   r=   rv   r   rw  r   poolerry  )r2   rb   add_pooling_layerrc   s      r.   rK   zAltRobertaModel.__init__  sE    
 	 .v6(02C&v.r3   Nrd   r   rG   rB   re   r   r&   c                    |du |duz  rt        d      | j                  ||||      }t        | j                  ||      } | j                  |fd|i|}|d   }| j
                  | j                  |      nd}	t        ||	      S )aK  
        Examples:

        ```python
        >>> from transformers import AutoTokenizer, AltRobertaModel

        >>> model = AltRobertaModel.from_pretrained("openai/alt_roberta-vit-base-patch32")
        >>> tokenizer = AutoTokenizer.from_pretrained("openai/alt_roberta-vit-base-patch32")

        >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")

        >>> outputs = model(**inputs)
        >>> last_hidden_state = outputs.last_hidden_state
        >>> pooled_output = outputs.pooler_output  # pooled (EOS token) states
        ```Nz:You must specify exactly one of input_ids or inputs_embeds)rd   rB   rG   re   )rb   re   r   r   r   r|  )r   rv   r	   rb   rw  r  r   )
r2   rd   r   rG   rB   re   r   r~  sequence_outputr   s
             r.   rw   zAltRobertaModel.forward  s    6 -t";<YZZ%)'	 ( 
 3;;')
 '$,,
)
 

 *!,8<8OO4UY)-'
 	
r3   )TNNNNN)r4   r5   r6   r   r:   ri  r  r   r   rp  rK   r   r   r   r8   r   r   r   r0   r   rw   r   r   s   @r.   r  r    s      *(-

   *..2.2,0-13
<<$&3
 t+3
 t+	3

 llT)3
 ||d*3
 +,3
 
+	+3
    3
r3   r  c                       e Zd ZU eed<   dZdZdZ f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j                  dz  dee   deez  fd              Z xZS )AltCLIPTextModelrb   r  rQ   robertac                 &   t         |   |       t        |d      | _        t	        j
                  |j                  |j                        | _        t	        j                  |j                  |j                        | _        | j                          y )NF)r  r@   )rJ   rK   r  r  rL   r   rO   project_dimtransformationrT   rU   pre_LNry  ra   s     r.   rK   zAltCLIPTextModel.__init__M  se     &vG ii(:(:F<N<NOll6#5#56;P;PQr3   Nrd   r   rG   rB   re   r   r&   c           	           | j                   d|||||d|}|d   }| j                  |      }| j                  |      }	|	dddf   }
t        |	|
|j                  |j
                        S )a+  
        Examples:

        ```python
        >>> from transformers import AutoProcessor, AltCLIPTextModel

        >>> model = AltCLIPTextModel.from_pretrained("BAAI/AltCLIP")
        >>> processor = AutoProcessor.from_pretrained("BAAI/AltCLIP")

        >>> texts = ["it's a cat", "it's a dog"]

        >>> inputs = processor(text=texts, padding=True, return_tensors="pt")

        >>> outputs = model(**inputs)
        >>> last_hidden_state = outputs.last_hidden_state
        >>> pooled_output = outputs.pooler_output  # pooled CLS states
        ```)rd   r   rG   rB   re   r   N)r   r}  r   rT  r;   )r  r  r  r   r   rT  )r2   rd   r   rG   rB   re   r   outputsr  projection_stater}  s              r.   rw   zAltCLIPTextModel.forwardT  s    : $,, 
))%'
 
 "!* ++o6  ..?(A.6.'!//))	
 	
r3   r  )r4   r5   r6   r   r:   ri  r  rh  rK   r   r   r8   r   r   r   r0   r   rw   r   r   s   @r.   r  r  G  s     *!  *..2.2,0-13
<<$&3
 t+3
 t+	3

 llT)3
 ||d*3
 +,3
 
8	83
  3
r3   r  logitsr&   c                     t         j                  j                  | t        j                  t        |       | j                              S )N)rk   )rL   r   cross_entropyr8   rZ   lenrk   )r  s    r.   contrastive_lossr    s/    ==&&vu||CKPVP]P]/^__r3   
similarityc                 R    t        |       }t        | j                        }||z   dz  S )Ng       @)r  T)r  caption_loss
image_losss      r.   image_text_contrastive_lossr    s*    #J/L!*,,/J:%,,r3   tensorc                     t        j                  | d      }t        j                  |dd      }t        j                  |d      }|S )z
    This method is equivalent to tensor.norm(p=2, dim=-1, keepdim=True) and used to make
    model `executorch` exportable. See issue https://github.com/pytorch/executorch/issues/3566
    r   rD   T)rh   keepdimr6  )r8   powsum)r  square_tensor
sum_tensornormed_tensors       r.   _get_vector_normr    s<    
 IIfa(M=b$?JIIj#.Mr3   c                       e Zd ZU eed<   def fdZee	 	 	 ddej                  dej                  dz  dej                  dz  dej                  dz  de
e   d	eez  fd
              Zee	 ddej                  dede
e   d	eez  f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j$                  dz  dedz  dede
e   d	eez  fd              Z xZS )r^  rb   c                    t         |   |       |j                  }|j                  }|j                  | _        |j
                  | _        |j                  | _        t        j                  | j                  j                        | _        t        j                  | j                  j                        | _        t        j                   | j                  | j                  d      | _        t        j                   | j                  | j                  d      | _        t        j&                  t)        j*                  | j                  j,                              | _        | j1                          y )NF)r'  )rJ   rK   text_configvision_configprojection_dimr  r`  rO   rb  r  _from_configrb   
text_modelrt  vision_modelrL   r   ra  r_  r*  r8   r  logit_scale_init_valuelogit_scalery  )r2   rb   r  r  rc   s       r.   rK   zAltCLIPModel.__init__  s     ((,,$33)55 - 9 9*778O8OP.;;DKK<U<UV!#4+@+@$BUBU\a!b!yy)<)<d>Q>QX]^<<T[[5W5W(XY 	r3   Nrd   r   rG   rB   r   r&   c                      | j                   d||||d|}|j                  dddddf   }| j                  |      |_        |S )a  
        Examples:

        ```python
        >>> import torch
        >>> from transformers import AutoProcessor, AltCLIPModel

        >>> model = AltCLIPModel.from_pretrained("BAAI/AltCLIP")
        >>> processor = AutoProcessor.from_pretrained("BAAI/AltCLIP")

        >>> inputs = processor(text=["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")
        >>> with torch.inference_mode():
        ...     text_features = model.get_text_features(**inputs)
        ```rd   r   rG   rB   Nr   r;   )r  r   r_  r}  )r2   rd   r   rG   rB   r   text_outputsr   s           r.   get_text_featureszAltCLIPModel.get_text_features  sa    0 4C4?? 4
))%	4

 4
 %66q!Qw?%)%9%9-%H"r3   rF  rE  c                 v     | j                   d||dd|}|j                  }| j                  |      |_        |S )ao  
        Examples:

        ```python
        >>> import torch
        >>> from transformers import AutoProcessor, AltCLIPModel
        >>> from transformers.image_utils import load_image

        >>> model = AltCLIPModel.from_pretrained("BAAI/AltCLIP")
        >>> processor = AutoProcessor.from_pretrained("BAAI/AltCLIP")

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = load_image(url)

        >>> inputs = processor(images=image, return_tensors="pt")
        >>> with torch.inference_mode():
        ...     image_features = model.get_image_features(**inputs)
        ```T)rF  rE  return_dictr;   )r  r}  ra  )r2   rF  rE  r   vision_outputsr   s         r.   get_image_featureszAltCLIPModel.get_image_features  sU    4 6GT5F5F 6
%%=6
 	6
 '44'+'='=m'L$r3   return_lossc           	      2    | j                   d||d|}	 | j                  d||||d|}
|	d   }| j                  |      }|
d   }| j                  |      }|t	        |      z  }|t	        |      z  }t        j                  ||j                         j                  |j                              }|| j                  j                         j                  |j                        z  }|j                         }d}|rt        |      }t        ||||||
|	      S )u  
        return_loss (`bool`, *optional*):
            Whether or not to return the contrastive loss.

        Examples:

        ```python
        >>> import torch
        >>> from transformers import AutoProcessor, AltCLIPModel
        >>> from transformers.image_utils import load_image

        >>> model = AltCLIPModel.from_pretrained("OFA-Sys/chinese-clip-vit-base-patch16")
        >>> processor = AutoProcessor.from_pretrained("OFA-Sys/chinese-clip-vit-base-patch16")

        >>> url = "https://clip-cn-beijing.oss-cn-beijing.aliyuncs.com/pokemon.jpeg"
        >>> image = load_image(url)

        >>> inputs = processor(text=["杰尼龟", "妙蛙种子", "小火龙", "皮卡丘"], images=image, return_tensors="pt", padding=True)

        >>> with torch.inference_mode():
        ...     outputs = model(**inputs)
        >>> logits_per_image = outputs.logits_per_image  # this is the image-text similarity score
        >>> probs = logits_per_image.softmax(dim=1)  # we can take the softmax to get the label probabilities
        ```)rF  rE  r  r   N)r   r    r!   r"   r#   r$   r%   r;   )r  r  ra  r_  r  r8   r   tro   rk   r  expr  r   )r2   rd   rF  r   rG   rB   r  rE  r   r  r  r#   r"   r!   r    r   s                   r.   rw   zAltCLIPModel.forward  sG   J +** 
%%=
 
 't 
))%	

 
 &a(--l;"1o**;7 $&6|&DD!$4[$AA  ,,{LNN4D4G4GHZHZ4[\)D,<,<,@,@,B,E,EkFXFX,YY*,,..?D-+#%* .
 	
r3   )NNNrM  )NNNNNNF)r4   r5   r6   r   r:   rK   r   r   r8   r   r   r   r0   r   r  r9   r  r  r   r   rw   r   r   s   @r.   r^  r^    s   } $  /3.2,0 <<  t+  t+	 
 llT)  +,  
+	+    D  */!''! #'! +,	!
 
+	+!  !F  .215.2.204#').L
##d*L
 ''$.L
 t+	L

 t+L
 &&-L
 D[L
 #'L
 +,L
 
	L
  L
r3   r^  )rO  rt  r  r^  )r   )Gcollections.abcr   dataclassesr   typingr   r8   torch.nnrL    r   rZ  activationsr   masking_utilsr	   modeling_layersr
   modeling_outputsr   r   r   modeling_utilsr   r   processing_utilsr   pytorch_utilsr   utilsr   r   r   r   r   utils.genericr   utils.output_capturingr   configuration_altclipr   r   r   r   Moduler=   r   floatr   r   r   r   r   r   r   r   r   r   r  r  r  r!  rO  rt  r  r  r  r  r  r^  __all__r;   r3   r.   <module>r     s  ( % !    & ! 6 9 t t F & 6 a a 7 5 X X 
_K _  _@g8299 g8b %II%<<% 
% <<	%
 LL4'% % %,4)bii 4)n299 ")) ,RYY ryy 0 B
		 
Dryy 7)ryy 7)t 4 B
RYY 
DPbii Pf 1/_ 1/ 1/h 
>
/ >

>
B L
, L
L
^B
- B
N`U\\ `ell `-ELL -U\\ -U\\ ell  l
) l
 l
^ _r3   