
    ^j*                        d Z ddl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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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.  e&j^                  e0      Z1 e%d      e G d de#                    Z2 G d dejf                        Z4 G d dejf                        Z5 G d dejf                        Z6d e5iZ7 G d! d"ejf                        Z8 G d# d$ejf                        Z9 G d% d&ejf                        Z: G d' d(e      Z; G d) d*ejf                        Z<e% G d+ d,e             Z= G d- d.ejf                        Z> G d/ d0ejf                        Z?	 dMd1ejf                  d2ej                  d3ej                  d4ej                  d5ej                  dz  d6eAd7eAfd8ZB G d9 d:ejf                        ZC G d; d<e      ZD G d= d>ejf                        ZE G d? d@ejf                        ZF e%dA       G dB dCe=             ZG G dD dEejf                        ZH e%dF       G dG dHe=             ZI e%dI       G dJ dKe=e             ZJg dLZKy)NzPyTorch GIT model.    N)Callable)	dataclass)nn   )initialization)ACT2FN)CacheDynamicCache)GenerationMixin)create_causal_mask)GradientCheckpointingLayer)BaseModelOutputBaseModelOutputWithPastBaseModelOutputWithPoolingCausalLMOutputWithPast)ALL_ATTENTION_FUNCTIONSPreTrainedModel)Unpack)apply_chunking_to_forward)ModelOutputTransformersKwargsauto_docstringlogging	torch_int)merge_with_config_defaults)capture_outputs   )	GitConfigGitVisionConfigz}
    Base class for vision model's outputs that also contains image embeddings of the pooling of the last hidden states.
    )custom_introc                       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
ej                  df   dz  ed<   dZe
ej                  df   dz  ed<   y)GitVisionModelOutputz
    image_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim)` *optional* returned when model is initialized with `with_projection=True`):
        The image embeddings obtained by applying the projection layer to the pooler_output.
    Nimage_embedslast_hidden_state.hidden_states
attentions)__name__
__module____qualname____doc__r#   torchFloatTensor__annotations__r$   r%   tupler&        o/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/git/modeling_git.pyr"   r"   6   sr    
 .2L%##d*126u((4/6:>M5**C/047>7;Je'',-4;r0   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dej                  f
d	Z
 xZS )GitEmbeddingsz;Construct the embeddings from word and position embeddings.c                    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       y )N)padding_idxepsposition_idsr   F
persistent)super__init__r   	Embedding
vocab_sizehidden_sizepad_token_idword_embeddingsmax_position_embeddingsposition_embeddings	LayerNormlayer_norm_epsDropouthidden_dropout_probdropoutregister_bufferr+   arangeexpandselfconfig	__class__s     r1   r>   zGitEmbeddings.__init__L   s    !||F,=,=v?Q?Q_e_r_rs#%<<0N0NPVPbPb#c f&8&8f>S>STzz&"<"<=ELL)G)GHOOPWXej 	 	
r0   N	input_idsr8   inputs_embedspast_key_values_lengthreturnc                 ,   ||j                         }n|j                         d d }|d   }|| j                  d d |||z   f   }|| j                  |      }n|}| j                  |      }||z  }| j	                  |      }| j                  |      }|S )Nr:   r   )sizer8   rC   rE   rF   rJ   )	rO   rR   r8   rS   rT   input_shape
seq_length
embeddingsrE   s	            r1   forwardzGitEmbeddings.forwardX   s      #..*K',,.s3K ^
,,Q0FVlIl0l-lmL --i8J&J"66|D))
^^J/
\\*-
r0   )NNNr   )r'   r(   r)   r*   r>   r+   
LongTensorr,   intTensorr[   __classcell__rQ   s   @r1   r3   r3   I   ss    E

 .20426&'##d* &&- ((4/	
 !$ 
r0   r3   c                        e Zd Zd	 fd	Z	 	 d
dej
                  dej                  dz  dedz  dee	   de
ej
                     f
dZ xZS )GitSelfAttentionNc                    t         |           |j                  |j                  z  dk7  r2t	        |d      s&t        d|j                   d|j                   d      || _        |-t        j                  d| j                  j                   d       |j                  | _        t        |j                  |j                  z        | _        | j                  | j                  z  | _        t        |j                  j                  |j                  j                   z  dz  d	z         | _        |j$                  | xj"                  |j$                  z  c_        t'        j(                  |j                  | j                        | _        t'        j(                  |j                  | j                        | _        t'        j(                  |j                  | j                        | _        t'        j0                  |j2                        | _        y )
Nr   embedding_sizezThe hidden size (z6) is not a multiple of the number of attention heads ()zInstantiating z without passing a `layer_idx` is not recommended and will lead to errors during the forward call if caching is used. Please make sure to provide a `layer_idx` when creating this class.   r   )r=   r>   rA   num_attention_headshasattr
ValueError	layer_idxloggerwarning_oncerQ   r'   r]   attention_head_sizeall_head_sizevision_config
image_size
patch_sizeimage_patch_tokensnum_image_with_embeddingr   LinearquerykeyvaluerH   attention_probs_dropout_probrJ   rO   rP   rj   rQ   s      r1   r>   zGitSelfAttention.__init__w   s    : ::a?PVXhHi#F$6$6#7 8 445Q8  # !8!8 9 :, , $*#=#= #&v'9'9F<V<V'V#W !558P8PP"%v';';'F'FI]I]IhIh'hmn&nqr&r"s**6##v'F'FF#YYv1143E3EF
99V//1C1CDYYv1143E3EF
zz&"E"EFr0   r%   attention_maskpast_key_valueskwargsrU   c                    |j                   d d }g |d| j                  }| j                  |      j                  |      j	                  dd      }| j                  |      j                  |      j	                  dd      }| j                  |      j                  |      j	                  dd      }	| |j                  ||	| j                        \  }}	t        j                  ||j	                  dd            }
|
t        j                  | j                        z  }
||
|z   }
t        j                  j                  |
d      }| j!                  |      }t        j                  ||	      }|j#                  dddd      j%                         }|j'                         d d | j(                  fz   }|j                  |      }||fS )Nr:   r   rf   dimr   r   )shaperm   ru   view	transposerv   rw   updaterj   r+   matmulmathsqrtr   
functionalsoftmaxrJ   permute
contiguousrW   rn   )rO   r%   rz   r{   r|   rX   hidden_shapequery_layer	key_layervalue_layerattention_scoresattention_probscontext_layernew_context_layer_shapes                 r1   r[   zGitSelfAttention.forward   s    $))#2.CCbC$*B*BCjj/44\BLLQPQRHH]+00>HHAN	jj/44\BLLQPQR&%4%;%;I{TXTbTb%c"I{ !<<Y5H5HR5PQ+dii8P8P.QQ%/.@ --//0@b/I ,,7_kB%--aAq9DDF"/"4"4"6s";t?Q?Q>S"S%**+BCo--r0   NNNr'   r(   r)   r>   r+   r^   r,   r	   r   r   r.   r[   r_   r`   s   @r1   rb   rb   v   sh    G> 48(,	%.||%. ))D0%. 	%.
 +,%. 
u||	%.r0   rb   c                   n     e Zd Z fdZdej
                  dej
                  dej
                  fdZ xZS )GitSelfOutputc                 (   t         |           t        j                  |j                  |j                        | _        t        j                  |j                  |j                        | _        t        j                  |j                        | _
        y Nr6   )r=   r>   r   rt   rA   denserF   rG   rH   rI   rJ   rN   s     r1   r>   zGitSelfOutput.__init__   s`    YYv1163E3EF
f&8&8f>S>STzz&"<"<=r0   r%   input_tensorrU   c                 r    | j                  |      }| j                  |      }| j                  ||z         }|S r   r   rJ   rF   rO   r%   r   s      r1   r[   zGitSelfOutput.forward   7    

=1]3}|'CDr0   r'   r(   r)   r>   r+   r^   r[   r_   r`   s   @r1   r   r      1    >U\\  RWR^R^ r0   r   eagerc                        e Zd Zd	 fd	Z	 	 d
dej
                  dej                  dz  dedz  dee	   de
ej
                     f
dZ xZS )GitAttentionNc                     t         |           t        |j                     ||      | _        t        |      | _        y )Nrj   )r=   r>   GIT_SELF_ATTENTION_CLASSES_attn_implementationrO   r   outputry   s      r1   r>   zGitAttention.__init__   s4    .v/J/JKF^gh	#F+r0   r%   rz   r{   r|   rU   c                 Z     | j                   |||fi |\  }}| j                  ||      }|S r   )rO   r   )rO   r%   rz   r{   r|   attn_output_attention_outputs           r1   r[   zGitAttention.forward   sD     #
 	
Q  ;;{MBr0   r   r   r   r`   s   @r1   r   r      sg    , 48(,	 ||  ))D0  	 
 +,  
u||	 r0   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )GitIntermediatec                    t         |           t        j                  |j                  |j
                        | _        t        |j                  t              rt        |j                     | _        y |j                  | _        y r   )r=   r>   r   rt   rA   intermediate_sizer   
isinstance
hidden_actstrr   intermediate_act_fnrN   s     r1   r>   zGitIntermediate.__init__   s]    YYv1163K3KL
f''-'-f.?.?'@D$'-'8'8D$r0   r%   rU   c                 J    | j                  |      }| j                  |      }|S r   )r   r   rO   r%   s     r1   r[   zGitIntermediate.forward   s&    

=100?r0   r   r`   s   @r1   r   r      s#    9U\\ ell r0   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 )	GitOutputc                 (   t         |           t        j                  |j                  |j
                        | _        t        j                  |j
                  |j                        | _        t        j                  |j                        | _        y r   )r=   r>   r   rt   r   rA   r   rF   rG   rH   rI   rJ   rN   s     r1   r>   zGitOutput.__init__   s`    YYv779K9KL
f&8&8f>S>STzz&"<"<=r0   r%   r   rU   c                 r    | j                  |      }| j                  |      }| j                  ||z         }|S r   r   r   s      r1   r[   zGitOutput.forward   r   r0   r   r`   s   @r1   r   r      r   r0   r   c                        e Zd Zd
 fd	Z	 	 ddej
                  dej                  dz  dedz  dee	   de
ej
                     f
dZd	 Z xZS )GitLayerNc                     t         |           |j                  | _        d| _        t	        ||      | _        t        |      | _        t        |      | _	        y )Nr   r   )
r=   r>   chunk_size_feed_forwardseq_len_dimr   	attentionr   intermediater   r   ry   s      r1   r>   zGitLayer.__init__  sK    '-'E'E$%f	B+F3'r0   r%   rz   r{   r|   rU   c                      | j                   ||fd|i|}t        | j                  | j                  | j                  |      }|S )Nr{   )r   r   feed_forward_chunkr   r   )rO   r%   rz   r{   r|   r   layer_outputs          r1   r[   zGitLayer.forward  s_     *4>>
 ,
 	
 1##T%A%A4CSCSUe
 r0   c                 L    | j                  |      }| j                  ||      }|S r   )r   r   )rO   r   intermediate_outputr   s       r1   r   zGitLayer.feed_forward_chunk!  s,    "//0@A{{#68HIr0   r   r   )r'   r(   r)   r>   r+   r^   r,   r	   r   r   r.   r[   r   r_   r`   s   @r1   r   r     sl    ( 48(,	|| ))D0 	
 +, 
u||	&r0   r   c                        e Zd Z fdZ	 	 	 d
dej
                  dej                  dz  dedz  dedz  de	e
   defd	Z xZS )
GitEncoderc           	          t         |           || _        t        j                  t        |j                        D cg c]  }t        ||       c}      | _        d| _	        y c c}w NF)
r=   r>   rP   r   
ModuleListrangenum_hidden_layersr   layergradient_checkpointing)rO   rP   irQ   s      r1   r>   zGitEncoder.__init__(  sP    ]]vG_G_A`#aAHVQ$7#ab
&+# $bs   A$Nr%   rz   r{   	use_cacher|   rU   c                 T    | j                   D ]  } ||||fi |} t        ||      S )Nr$   r{   )r   r   )rO   r%   rz   r{   r   r|   layer_modules          r1   r[   zGitEncoder.forward.  sI     !JJ 	L( 	M	 '++
 	
r0   )NNN)r'   r(   r)   r>   r+   r^   r,   r	   boolr   r   r   r[   r_   r`   s   @r1   r   r   '  so    , 48(,!%
||
 ))D0
 	

 $;
 +,
 
!
r0   r   c                   ^     e Zd ZU eed<   dZdZdZ ej                          fd       Z
 xZS )GitPreTrainedModelrP   git)imagetextTc                 4   t         |   |       t        |t              rt	        j
                  |j                  d| j                  j                         t	        j
                  |j                  j                  | j                  j                         t	        j
                  |j                  j                  | j                  j                         t	        j                  |j                  t        j                  |j                  j                   d         j#                  d             t        |t$              rZt	        j                  |j                  t        j                  |j                  j                   d         j#                  d             yy)zInitialize the weights        )meanstd)r   r:   r9   N)r=   _init_weightsr   GitVisionEmbeddingsinitnormal_class_embeddingrP   initializer_rangepatch_embeddingweightposition_embeddingcopy_r8   r+   rL   r   rM   r3   )rO   modulerQ   s     r1   r   z GitPreTrainedModel._init_weightsK  s
    	f%f12LL//ct{{?\?\]LL//66DKK<Y<YZLL2299t{{?\?\]JJv**ELL9L9L9R9RSU9V,W,^,^_f,ghfm,JJv**ELL9L9L9R9RSU9V,W,^,^_f,gh -r0   )r'   r(   r)   r   r-   base_model_prefixinput_modalitiessupports_gradient_checkpointingr+   no_gradr   r_   r`   s   @r1   r   r   D  s7    (&*#U]]_	i 	ir0   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 )r   rP   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biasrf   r   r8   r9   r;   )r=   r>   rP   rA   	embed_dimrp   rq   r   	Parameterr+   randnr   Conv2dnum_channelsr   num_patchesnum_positionsr?   r   rK   rL   rM   rN   s     r1   r>   zGitVisionEmbeddings.__init__Z  s	   ++ ++ ++!||EKK,GH!yy++?? 
 !OOt>1D!--1"$,,t/A/A4>>"R^U\\$:L:L-M-T-TU\-]jopr0   rZ   heightwidthrU   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   Nr:   g      ?r   rf   bicubicF)rW   modealign_cornersr   )r   r   r   	unsqueezer+   jit
is_tracingr8   rq   r   reshaper   r   r   interpolater   cat)rO   rZ   r   r   r   r   r   class_pos_embedpatch_pos_embedr   
new_height	new_widthsqrt_num_positionss                r1   interpolate_pos_encodingz,GitVisionEmbeddings.interpolate_pos_encodingp  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Cr0   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 ().dtyperf   r   r:   r   )r   rp   ri   r   r   r  toflattenr   r   rM   r+   r	  r  r   r8   )rO   r  r  
batch_sizer   r   r   target_dtypepatch_embedsclass_embedsrZ   s              r1   r[   zGitVisionEmbeddings.forward  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r0   )F)r'   r(   r)   r   r>   r+   r^   r]   r  r,   r[   r_   r`   s   @r1   r   r   Y  sd    q q,'D5<< 'D 'DUX 'D]b]i]i 'DRE$5$5 Z_ZfZf r0   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )GitVisionMLPc                    t         |           || _        t        |j                     | _        t        j                  |j                  |j                        | _
        t        j                  |j                  |j                        | _        y r   )r=   r>   rP   r   r   activation_fnr   rt   rA   r   fc1fc2rN   s     r1   r>   zGitVisionMLP.__init__  sd    #F$5$5699V//1I1IJ99V55v7I7IJr0   r%   rU   c                 l    | j                  |      }| j                  |      }| j                  |      }|S r   )r   r  r!  r   s     r1   r[   zGitVisionMLP.forward  s4    /**=9/r0   r   r`   s   @r1   r  r    s$    KU\\ ell r0   r  r   ru   rv   rw   rz   scalingrJ   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 )Nr:   r~   )r   r  )ptrainingr   rf   )r+   r   r   r   r   r   float32r  r  rJ   r&  r   )
r   ru   rv   rw   rz   r#  rJ   r|   attn_weightsr   s
             r1   eager_attention_forwardr)    s     <<s}}R'<=GL!#n4==((2U]](SVVW\WbWbcL==((6??([L,,|U3K''1-88:K$$r0   c                        e Zd 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 )
GitVisionAttentionz=Multi-headed attention from 'Attention Is All You Need' paperc                    t         |           || _        |j                  | _        |j
                  | _        | j                  | j                  z  | _        | j                  | j                  z  | j                  k7  r&t        d| j                   d| j                   d      | j                  dz  | _	        |j                  | _        d| _        t        j                  | j                  | j                        | _        t        j                  | j                  | j                        | _        t        j                  | j                  | j                        | _        t        j                  | j                  | j                        | _        y )Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: r  g      F)r=   r>   rP   rA   r   rg   	num_headshead_dimri   scaleattention_dropoutrJ   	is_causalr   rt   k_projv_projq_projout_projrN   s     r1   r>   zGitVisionAttention.__init__  s   ++33$..8==4>>)T^^;MdnnM] ^NN#2'  ]]D(
//ii?ii?ii?		$..$..Ar0   Nr%   rz   r|   rU   c                    |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                  | j                  | j                  sdn| j                  d|\  }
} |
j                   g |d j#                         }
| j%                  |
      }
|
|fS )z#Input shape: Batch x Time x ChannelNr:   r   rf   r   )r1  r#  rJ   )r   r.  r4  r   r   r2  r3  r   get_interfacerP   r   r)  r1  r/  r&  rJ   r  r   r5  )rO   r%   rz   r|   rX   r   querieskeysvaluesattention_interfacer   r(  s               r1   r[   zGitVisionAttention.forward  sM    $))#2.88b8$--8++m,11,?II!QO{{=)..|<FFq!L]+00>HHAN(?(M(MKK,,.E)
 %8
%
 nnJJ#}}C$,,
%
 
%
!\ *k));;;;FFHmmK0L((r0   r   )r'   r(   r)   r*   r>   r+   r^   r   r   r.   r[   r_   r`   s   @r1   r+  r+    sf    GB. /3)||) t+) +,	)
 
u||U\\D00	1)r0   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 )GitVisionEncoderLayerrP   c                 D   t         |           |j                  | _        t	        |      | _        t        j                  | j                  |j                        | _	        t        |      | _        t        j                  | j                  |j                        | _        y r   )r=   r>   rA   r   r+  	self_attnr   rF   rG   layer_norm1r  mlplayer_norm2rN   s     r1   r>   zGitVisionEncoderLayer.__init__  sm    +++F3<<F<Q<QR'<<F<Q<QRr0   r%   rz   r|   rU   c                     |}| j                  |      } | j                  d||d|\  }}||z   }|}| j                  |      }| j                  |      }||z   }|S )N)r%   rz   r/   )r@  r?  rB  rA  )rO   r%   rz   r|   residualr   s         r1   r[   zGitVisionEncoderLayer.forward  s     !((7)4>> 
')
 
q
 !=0 ((7/ =0r0   )r'   r(   r)   r   r>   r+   r^   r   r   r,   r[   r_   r`   s   @r1   r=  r=    sP    S S||  +,	
 
		r0   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 )
GitVisionEncoderz
    Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a
    [`GitVisionEncoderLayer`].

    Args:
        config: GitVisionConfig
    rP   c                     t         |           || _        t        j                  t        |j                        D cg c]  }t        |       c}      | _        d| _	        y c c}w r   )
r=   r>   rP   r   r   r   r   r=  layersr   )rO   rP   r   rQ   s      r1   r>   zGitVisionEncoder.__init__8  sP    mmERXRjRjLk$lq%:6%B$lm&+# %ms   A#Nrz   r|   rU   c                 T    |}| j                   D ]  } |||fi |} t        |      S )Nr$   )rH  r   )rO   rS   rz   r|   r%   encoder_layers         r1   r[   zGitVisionEncoder.forward>  sH     &![[ 	M) M	 +
 	
r0   r   )r'   r(   r)   r*   r   r>   r+   r^   r   r   r   r[   r_   r`   s   @r1   rF  rF  /  sK    , , /3
 t+
 +,	

 

r0   rF  c            
       r     e Zd Zdef fdZe	 	 d	dej                  dz  dedz  de	e
   defd       Z xZS )
GitVisionTransformerrP   c                     t         |           || _        |j                  }t	        |      | _        t        j                  ||j                        | _	        t        |      | _        t        j                  ||j                        | _        y r   )r=   r>   rP   rA   r   rZ   r   rF   rG   pre_layrnormrF  encoderpost_layernorm)rO   rP   r   rQ   s      r1   r>   zGitVisionTransformer.__init__R  sj    &&	-f5LL8M8MN'/ ll9&:O:OPr0   Nr  r  r|   rU   c                     |t        d      | j                  ||      }| j                  |      } | j                  dd|i|}|j                  }| j                  |      }t        |      S )Nz You have to specify pixel_valuesr  rS   rJ  r/   )ri   rZ   rO  rP  r$   rQ  r   )rO   r  r  r|   r%   encoder_outputsr$   s          r1   r[   zGitVisionTransformer.forward\  s     ?@@Ogh))-8&$,, 
'


 ,== //0AB/
 	
r0   r   )r'   r(   r)   r   r>   r   r+   r,   r   r   r   r   r[   r_   r`   s   @r1   rM  rM  Q  sh    Q Q  2605
''$.
 #'+
 +,	

 

 
r0   rM  zY
    The vision model from CLIP, used in GIT, without any head or projection on top.
    c                        e Zd ZU eed<   dZdZeedZ	def fdZ
dej                  fdZe ed	      e	 	 ddej$                  d
z  dedee   deez  fd                     Z xZS )GitVisionModelrP   r  )r   r%   r&   c                 d    t         |   |       t        |      | _        | j	                          y r   )r=   r>   rM  vision_model	post_initrN   s     r1   r>   zGitVisionModel.__init__  s'     08r0   rU   c                 B    | j                   j                  j                  S r   )rY  rZ   r   rO   s    r1   get_input_embeddingsz#GitVisionModel.get_input_embeddings  s      ++;;;r0   F)tie_last_hidden_statesNr  r|   c                 ,     | j                   d||d|S )a  
        Examples:

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

        >>> processor = AutoProcessor.from_pretrained("microsoft/git-base")
        >>> model = GitVisionModel.from_pretrained("microsoft/git-base")

        >>> 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
        ```)r  r  r/   )rY  )rO   r  r  r|   s       r1   r[   zGitVisionModel.forward  s.    < !t   
%%=
 
 	
r0   r   )r'   r(   r)   r   r-   main_input_namer   r=  r+  _can_record_outputsr>   r   Moduler]  r   r   r   r+   r,   r   r   r   r.   r   r[   r_   r`   s   @r1   rV  rV  w  s     $O!.(
 <bii <  E2 26).
''$.
 #'
 +,	

 
	 
  3  
r0   rV  c                   \     e Zd Zdef fdZdej                  dej                  fdZ xZS )GitProjectionrP   c                 0   t         |           || _        t        j                  t        j
                  |j                  j                  |j                        t        j                  |j                  |j                  j                              | _
        y r   )r=   r>   rP   r   
Sequentialrt   ro   rA   rF   rG   visual_projectionrN   s     r1   r>   zGitProjection.__init__  sf    !#IIf**668J8JKLL++1E1E1T1TU"
r0   rZ   rU   c                 $    | j                  |      S r   )rg  )rO   rZ   s     r1   r[   zGitProjection.forward  s    %%j11r0   )	r'   r(   r)   r   r>   r+   r^   r[   r_   r`   s   @r1   rd  rd    s*    
y 
2%,, 25<< 2r0   rd  zy
    The bare GIT Model transformer consisting of a CLIP image encoder and text decoder outputting raw hidden-states
    c                   H    e Zd ZeedZ fdZd Z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dz  dedz  dedee   deej                     ez  fd                     Z xZS )GitModelrW  c                 l   t         |          | _        t              | _        t        j                        | _        t              | _	        t              | _        j                  6t        j                  fdt        j                        D              | _        | j#                          y )Nc              3      K   | ]B  }t        j                  t        j                  d d j                  j
                               D yw)r   N)r   r   r+   zerosro   rA   ).0r   rP   s     r1   	<genexpr>z$GitModel.__init__.<locals>.<genexpr>  s;      ; U[[Av/C/C/O/OPQ;s   AA)r=   r>   rP   r3   rZ   rV  ro   image_encoderr   rP  rd  rg  rs   r   ParameterListr   img_temporal_embeddingrZ  rN   s    `r1   r>   zGitModel.__init__  s     '/+F,@,@A!&)!.v!6**6*,*:*: ;v>>?; +D' 	r0   c                 .    | j                   j                  S r   rZ   rC   r\  s    r1   r]  zGitModel.get_input_embeddings  s    ...r0   c                 &    || j                   _        y r   rt  )rO   rw   s     r1   set_input_embeddingszGitModel.set_input_embeddings  s    */'r0   NrR   rz   r8   r  rS   r{   r   r  r|   rU   c	           	      "   |du |duz  rt        d      |r|t        | j                        }d}
|0t        |t              s|j                         n|j                         }
|||j                  d   dk(  r||
z   }| j                  ||||
      }t        j                  |t        j                        d   }||j                  d	k(  r| j                  ||
      j                  }n|j                  dk(  rg }t        |j                  d         D ]O  }| j                  |dd|ddddf   |
      j                  }|| j                  |   z  }|j!                  |       Q t        j"                  |d      }nt        d      | j%                  |      }|j'                  |j)                  d      |j)                  d      z  dd      }t        j"                  ||fd      }t        j*                  ||j,                        d   }t        j"                  ||gd      }|t        j"                  t        j*                  ||j,                        |gd      }n{|y|j                  d   dk(  rgt        j.                  |j                  d   |
|j                  d   z
  dz   f|j,                  |j0                        }t        j"                  ||gd      }t        j2                  g |j)                         dd d|j0                        }|t        j4                  |dk(  dd      }| j                  j7                         |||||d}t9        di |}|} | j:                  |f|||d|	}t=        |j                  |j>                        S )a   
        Examples:

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

        >>> processor = AutoProcessor.from_pretrained("microsoft/git-base")
        >>> model = AutoModel.from_pretrained("microsoft/git-base")

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

        >>> text = "this is an image of two cats"

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

        >>> outputs = model(**inputs)
        >>> last_hidden_state = outputs.last_hidden_state
        ```Nz:You must specify exactly one of input_ids or inputs_embeds)rP   r   r   )rR   r8   rS   rT   r  ).r      rS     r   z#pixel_values must be of rank 4 or 5r:   )r  device)rz  )rP   rS   rz   r{   r8   block_sequence_ids)rz   r{   r   r   r/   ) ri   r
   rP   r   r	   get_seq_lengthr   rZ   r+   
zeros_liker]   ndimrp  r$   r   rr  appendr	  rg  repeatrW   	ones_liker  onesrz  fullwhereget_text_configr   rP  r   r{   )rO   rR   rz   r8   r  rS   r{   r   r  r|   rT   embedding_outputtoken_type_idsvisual_features	frame_idxvisual_features_frameprojected_visual_featuresimage_token_type_idsextended_attention_mask	group_idsmask_kwargscausal_maskr%   rT  s                           r1   r[   zGitModel.forward  s   L -t";<YZZ0*$++>O "#& "/59  ..0$335 # O$?IOOTUDVZ[D['*@@L??%'#9	 + 
 ))*:%))LVT#  A%"&"4"4 ;S #5 ###   ""a'"$!&|'9'9!'<!= BI,0,>,>$Q	1a%78Sk -? -'' * *T-H-H-SS)#**+@AB #())O"C !!FGG(,(>(>(O% )B(H(H %%a(,E,J,J1,MMqRS)%
  %yy*CEU)V\]^#(??3LTbThTh#ijp#q "YY(<n'MSUVN)!&__%9AUAUVXfgmo" (Y__Q-?1-D ',jj%%a(*@>CWCWXYCZ*Z]^*^_$**%,,'#
 #YY(?'PVXYN JJ>!1!6!6!8"!=>K[KbKbc	%Na$7B?I kk113-,.("+
 )7;7(3?4<<4
&+	4

 4
 '-??+;;
 	
r0   )NNNNNNNF)r'   r(   r)   r   rb   ra  r>   r]  rv  r   r   r   r+   r^   r	   r   r   r   r.   r   r[   r_   r`   s   @r1   rj  rj    s    "&
&/0   *..2,0,0-1(,!%).N
<<$&N
 t+N
 llT)	N

 llT)N
 ||d*N
 N
 $;N
 #'N
 +,N
 
u||	9	9N
    N
r0   rj  z`
    GIT Model with a `language modeling` head on top for autoregressive language modeling.
    c                       e Zd ZddiZ fdZd Z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
j                  dz  dedz  dedz  dedee
j                  z  dee   dee
j                     ez  fd                     Z	 	 	 	 	 d fd	Z xZS )GitForCausalLMzoutput.weightz%git.embeddings.word_embeddings.weightc                     t         |   |       t        |      | _        t	        j
                  |j                  |j                        | _        | j                          y r   )
r=   r>   rj  r   r   rt   rA   r@   r   rZ  rN   s     r1   r>   zGitForCausalLM.__init__  sF     F#ii 2 2F4E4EF 	r0   c                     | j                   S r   r   r\  s    r1   get_output_embeddingsz$GitForCausalLM.get_output_embeddings  s    {{r0   c                     || _         y r   r  )rO   new_embeddingss     r1   set_output_embeddingsz$GitForCausalLM.set_output_embeddings  s	    $r0   NrR   rz   r8   r  rS   labelsr{   r   r  logits_to_keepr|   rU   c                    |d} | j                   |f|||||||	d|}|j                  }t        |
t              rt	        |
 d      n|
}| j                  |dd|ddf         }d}|| j                   j                  j                  d   j                  j                  j                  }|dd|dddf   j                         }|ddddf   j                         } | j                  |j                  d| j                  j                        |j                  d      fd| j                  j                  i|}t!        |||j"                  |j$                  |j&                        S )	a0  
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the left-to-right language modeling loss (next word prediction). Indices should be in
            `[-100, 0, ..., config.vocab_size]` (see `input_ids` docstring) Tokens with indices set to `-100` are
            ignored (masked), the loss is only computed for the tokens with labels n `[0, ..., config.vocab_size]`

        Examples:

        Image captioning example:

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

        >>> processor = AutoProcessor.from_pretrained("microsoft/git-base-coco")
        >>> model = AutoModelForCausalLM.from_pretrained("microsoft/git-base-coco")

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

        >>> pixel_values = processor(images=image, return_tensors="pt").pixel_values

        >>> generated_ids = model.generate(pixel_values=pixel_values, max_length=50)
        >>> generated_caption = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
        >>> print(generated_caption)
        two cats sleeping on a pink blanket next to remotes.
        ```

        Visual question answering (VQA) example:

        ```python
        >>> from transformers import AutoProcessor, AutoModelForCausalLM
        >>> from huggingface_hub import hf_hub_download
        >>> from PIL import Image

        >>> processor = AutoProcessor.from_pretrained("microsoft/git-base-textvqa")
        >>> model = AutoModelForCausalLM.from_pretrained("microsoft/git-base-textvqa")

        >>> file_path = hf_hub_download(repo_id="nielsr/textvqa-sample", filename="bus.png", repo_type="dataset")
        >>> image = Image.open(file_path).convert("RGB")

        >>> pixel_values = processor(images=image, return_tensors="pt").pixel_values

        >>> question = "what does the front of the bus say at the top?"

        >>> input_ids = processor(text=question, add_special_tokens=False).input_ids
        >>> input_ids = [processor.tokenizer.cls_token_id] + input_ids
        >>> input_ids = torch.tensor(input_ids).unsqueeze(0)

        >>> generated_ids = model.generate(pixel_values=pixel_values, input_ids=input_ids, max_length=50)
        >>> print(processor.batch_decode(generated_ids, skip_special_tokens=True))
        ['what does the front of the bus say at the top? special']
        ```

        Video captioning example:

        ```python
        >>> import av
        >>> import numpy as np
        >>> from PIL import Image
        >>> from huggingface_hub import hf_hub_download
        >>> from transformers import AutoProcessor, AutoModelForCausalLM

        >>> processor = AutoProcessor.from_pretrained("microsoft/git-base-vatex")
        >>> model = AutoModelForCausalLM.from_pretrained("microsoft/git-base-vatex")

        >>> # set seed for reproducibility
        >>> np.random.seed(45)


        >>> def read_video_pyav(container, indices):
        ...     '''
        ...     Decode the video with PyAV decoder.
        ...     Args:
        ...         container (`av.container.input.InputContainer`): PyAV container.
        ...         indices (`list[int]`): List of frame indices to decode.
        ...     Returns:
        ...         result (np.ndarray): np array of decoded frames of shape (num_frames, height, width, 3).
        ...     '''
        ...     frames = []
        ...     container.seek(0)
        ...     start_index = indices[0]
        ...     end_index = indices[-1]
        ...     for i, frame in enumerate(container.decode(video=0)):
        ...         if i > end_index:
        ...             break
        ...         if i >= start_index and i in indices:
        ...             frames.append(frame)
        ...     return np.stack([x.to_ndarray(format="rgb24") for x in frames])


        >>> def sample_frame_indices(clip_len, frame_sample_rate, seg_len):
        ...     '''
        ...     Sample a given number of frame indices from the video.
        ...     Args:
        ...         clip_len (`int`): Total number of frames to sample.
        ...         frame_sample_rate (`int`): Sample every n-th frame.
        ...         seg_len (`int`): Maximum allowed index of sample's last frame.
        ...     Returns:
        ...         indices (`list[int]`): List of sampled frame indices
        ...     '''
        ...     converted_len = int(clip_len * frame_sample_rate)
        ...     end_idx = np.random.randint(converted_len, seg_len)
        ...     start_idx = end_idx - converted_len
        ...     indices = np.linspace(start_idx, end_idx, num=clip_len)
        ...     indices = np.clip(indices, start_idx, end_idx - 1).astype(np.int64)
        ...     return indices


        >>> # load video
        >>> file_path = hf_hub_download(
        ...     repo_id="nielsr/video-demo", filename="eating_spaghetti.mp4", repo_type="dataset"
        ... )
        >>> container = av.open(file_path)

        >>> # sample frames
        >>> num_frames = model.config.num_image_with_embedding
        >>> indices = sample_frame_indices(
        ...     clip_len=num_frames, frame_sample_rate=4, seg_len=container.streams.video[0].frames
        ... )
        >>> frames = read_video_pyav(container, indices)

        >>> pixel_values = processor(images=list(frames), return_tensors="pt").pixel_values

        >>> generated_ids = model.generate(pixel_values=pixel_values, max_length=50)

        >>> print("Generated caption:", processor.batch_decode(generated_ids, skip_special_tokens=True))
        Generated caption: ['a woman is sitting at a table and she is talking about the food she is holding.']
        ```
        NF)rz   r8   r  rS   r{   r   r  r   r:   r   r@   )losslogitsr{   r%   r&   )r   r$   r   r]   slicer   rP  r   r   rO   rr   r   loss_functionr   rP   r@   r   r{   r%   r&   )rO   rR   rz   r8   r  rS   r  r{   r   r  r  r|   outputsr%   slice_indicesr  r  num_image_tokensshifted_logitss                      r1   r[   zGitForCausalLM.forward  s~   l I+3488
,
)%%'+%=
,
 
,
  118B>SV8W~ot4]k]1mQ+>?@#xx//55a8BBGGZZ#A'7':A$=>IIKNAqrE]--/F%4%%##B(>(>?B  ;;11 	D &#33!//))
 	
r0   c                 D    t        	|   |f||||d|}|s|s||d<   |S )N)r{   rz   r   is_first_iterationr  )r=   prepare_inputs_for_generation)
rO   rR   r{   r  rz   r   r  r|   model_inputsrQ   s
            r1   r  z,GitForCausalLM.prepare_inputs_for_generationN  sI     w<
+)1
 
 Y+7L(r0   )
NNNNNNNNFr   )NNNNF)r'   r(   r)   _tied_weights_keysr>   r  r  r   r   r   r+   r^   r	   r   r]   r   r   r.   r   r[   r  r_   r`   s   @r1   r  r  x  s`    *+RS%   *..2,0,0-1&*(,!%).-.z
<<$&z
 t+z
 llT)	z

 llT)z
 ||d*z
 t#z
 z
 $;z
 #'z
 ell*z
 +,z
 
u||	5	5z
    z
~   r0   r  )r  rj  r   rV  )r   )Lr*   r   collections.abcr   dataclassesr   r+   r    r   r   activationsr   cache_utilsr	   r
   
generationr   masking_utilsr   modeling_layersr   modeling_outputsr   r   r   r   modeling_utilsr   r   processing_utilsr   pytorch_utilsr   utilsr   r   r   r   r   utils.genericr   utils.output_capturingr   configuration_gitr   r   
get_loggerr'   rk   r"   rb  r3   rb   r   r   r   r   r   r   r   r   r   r  r^   floatr)  r+  r=  rF  rM  rV  rd  rj  r  __all__r/   r0   r1   <module>r     s     $ !   & ! . ) / 9  G & 6  8 5 9 
		H	% 
 	<; 	< 	<*BII *ZB.ryy B.LBII   
 299  0bii  		 ) D
 
: i i i(P")) Pf299 . %II%<<% 
% <<	%
 LL4'% % %.6) 6)t6 D
ryy 
D#
299 #
L 
4
' 4

4
n
2BII 
2 
p
! p

p
f 
i' i
iX Qr0   