
    ^jK                     *   d dl mZmZ d dlmZ d dl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 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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+  e#d      e G d de                    Z, G d dejZ                        Z. G d dejZ                        Z/ G d dejZ                        Z0	 	 dOdejZ                  dej                  d ej                  d!ej                  d"ej                  dz  d#e1dz  d$e1d%ee"   fd&Z2 G d' d(ejZ                        Z3 G d) d*ejZ                        Z4 G d+ d,ejZ                        Z5 G d- d.e      Z6e# G d/ d0e             Z7e# G d1 d2e7             Z8 G d3 d4ejZ                        Z9 e#d5       G d6 d7e7             Z: e#d8       G d9 d:e7             Z; G d; d<ejZ                        Z< G d= d>ejZ                        Z= G d? d@ejZ                        Z> G dA dBejZ                        Z? G dC dDejZ                        Z@ G dE dFejZ                        ZA G dG dHejZ                        ZBe# G dI dJe7             ZC e#dK       G dL dMee7             ZDg dNZEy)P    )CallableIterable)	dataclassN)Tensornn   )initialization)ACT2FN)BackboneMixinfilter_output_hidden_states)create_bidirectional_mask)GradientCheckpointingLayer)BackboneOutputBaseModelOutputWithPoolingImageClassifierOutputMaskedLMOutputSemanticSegmenterOutput)ALL_ATTENTION_FUNCTIONSPreTrainedModel)Unpack)#compile_compatible_method_lru_cache)TransformersKwargsauto_docstring	torch_int)can_return_tuplemerge_with_config_defaults)capture_outputs   )
BeitConfigz-
    Class for outputs of [`BeitModel`].
    )custom_introc                       e Zd ZdZy)BeitModelOutputWithPoolingaF  
    pooler_output (`torch.FloatTensor` of shape `(batch_size, hidden_size)`):
        Average of the last layer hidden states of the patch tokens (excluding the *[CLS]* token) if
        *config.use_mean_pooling* is set to True. If set to False, then the final hidden state of the *[CLS]* token
        will be returned.
    N)__name__
__module____qualname____doc__     q/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/beit/modeling_beit.pyr"   r"   0   s    r(   r"   c                   `     e Zd ZdZdef fdZdej                  dej                  fdZ xZ	S )BeitPatchEmbeddingsz
    This class turns `pixel_values` of shape `(batch_size, num_channels, height, width)` into the initial
    `hidden_states` (patch embeddings) of shape `(batch_size, seq_length, hidden_size)` to be consumed by a
    Transformer.
    configc                    t         |           |j                  }|j                  }t	        |t
              r|n||f}t	        |t
              r|n||f}|d   |d   z  |d   |d   z  z  | _        || _        || _        |j                  | _        t        j                  |j                  |j                  ||      | _        y )Nr   r   kernel_sizestride)super__init__
image_size
patch_size
isinstancer   num_patchesnum_channelsr   Conv2dhidden_size
projection)selfr,   r3   r4   	__class__s       r)   r2   zBeitPatchEmbeddings.__init__F   s    &&
&&
#-j(#CZ*V`Ia
#-j(#CZ*V`Ia
&qMZ]:z!}PZ[\P]?]^$$"//))F$7$79K9KYclvwr(   pixel_valuesreturnc                     |j                   d   }|| j                  k7  rt        d| j                   d| d      | j                  |      j	                  d      j                  dd      S )Nr   zoMake sure that the channel dimension of the pixel values match with the one set in the configuration. Expected z	 but got .   )shaper7   
ValueErrorr:   flatten	transpose)r;   r=   r7   s      r)   forwardzBeitPatchEmbeddings.forwardS   su    #))!,4,,,!../yaI  |,44Q7AA!QGGr(   )
r#   r$   r%   r&   r   r2   torchr   rF   __classcell__r<   s   @r)   r+   r+   ?   s4    xz xHELL HU\\ Hr(   r+   c                        e Zd ZdZdeddf fdZdej                  dededej                  fd	Z		 dd
ej                  dej                  dz  dej                  fdZ xZS )BeitEmbeddingszb
    Construct the CLS token, position and patch embeddings. Optionally, also the mask token.
    r,   r>   Nc                 l   t         |           t        j                  t	        j
                  dd|j                              | _        |j                  r4t        j                  t	        j
                  dd|j                              nd | _	        t        |      | _        |j                  | _        | j                  j                  }|j                  r7t        j                  t	        j
                  d|dz   |j                              nd | _        t        j                   |j"                        | _        y )Nr   )r1   r2   r   	ParameterrG   zerosr9   	cls_tokenuse_mask_token
mask_tokenr+   patch_embeddingsr4   r6    use_absolute_position_embeddingsposition_embeddingsDropouthidden_dropout_probdropout)r;   r,   r6   r<   s      r)   r2   zBeitEmbeddings.__init__b   s    ekk!Q8J8J&KLQWQfQf",,u{{1a9K9K'LMlp 3F ; ++++77 66 LLQa9K9KLM 	 
 zz&"<"<=r(   
embeddingsheightwidthc                    |j                   d   dz
  }| j                  j                   d   dz
  }t        j                  j	                         s||k(  r||k(  r| j                  S | j                  ddddf   }| j                  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   Ng      ?r   r   rA   bicubicFsizemodealign_cornersdim)rB   rT   rG   jit
is_tracingr4   r   reshapepermuter   
functionalinterpolateviewcat)r;   rX   rY   rZ   r6   num_positionsclass_pos_embedpatch_pos_embedrc   
new_height	new_widthsqrt_num_positionss               r)   interpolate_pos_encodingz'BeitEmbeddings.interpolate_pos_encodingq   s`    !&&q)A-0066q9A= yy##%+*F6UZ?+++221bqb59221ab59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Cr(   r=   bool_masked_posc                    |j                   \  }}}}| j                  |      }|j                         \  }}}|K| j                  j	                  ||d      }	|j                  d      j                  |	      }
|d|
z
  z  |	|
z  z   }| j                  j	                  |dd      }t        j                  ||fd      }| j                  || j                  |||      z   }| j                  |      }|S Nr\   r   rb   )rB   rR   r_   rQ   expand	unsqueezetype_asrO   rG   rk   rT   rr   rW   )r;   r=   rs   _rY   rZ   rX   
batch_sizeseq_lenmask_tokensmask
cls_tokenss               r)   rF   zBeitEmbeddings.forward   s    
 +001fe**<8
!+!2
GQ&//00WbIK",,R088ED#q4x0;3EEJ^^**:r2>
YY
J7Q?
##/#d&C&CJPVX]&^^J\\*-
r(   N)r#   r$   r%   r&   r   r2   rG   r   intrr   
BoolTensorrF   rH   rI   s   @r)   rK   rK   ]   s    >z >d >&D5<< &D &DUX &D]b]i]i &DV 48ll ))D0 
	r(   rK   c                        e Zd Zdeddf fdZe ed      deeef   de	j                  fd              Zdd	ede	j                  fd
Z xZS )BeitRelativePositionBiasr,   r>   Nc                    t         |           |j                  }t        |t        t
        f      s||f}|d   |j                  z  |d   |j                  z  f| _        d| j                  d   z  dz
  d| j                  d   z  dz
  z  dz   | _        t        j                  t        j                  | j                  |j                              | _        y Nr   r   rA   r   )r1   r2   r3   r5   tuplelistr4   window_sizenum_relative_distancer   rM   rG   rN   num_attention_headsrelative_position_bias_table)r;   r,   r3   r<   s      r)   r2   z!BeitRelativePositionBias.__init__   s    &&
*udm4$j1J&qMV->->>
1QWQbQb@bc&'$*:*:1*=&=&Aa$JZJZ[\J]F]`aFa%bef%f",.LLKK22F4N4NO-
)r(   
   )maxsizer   c                    d| d   z  dz
  d| d   z  dz
  z  dz   }| d   | d   z  }t        j                  t        j                  t        j                  t        j                  | d         t        j                  | d         d            d      }|dddddf   |dddddf   z
  j                  ddd      j                         }|dddddfxx   | d   dz
  z  cc<   |dddddfxx   | d   dz
  z  cc<   |dddddfxx   d| d   z  dz
  z  cc<   t        j                  |dz   fdz  |j                  	      }|j                  d
      |ddddf<   |dz
  |dddf<   |dz
  |dddf<   |dz
  |d<   |S )z
        This method creates the relative position index, modified to support arbitrary window sizes,
        as introduced in [MiDaS v3.1](https://huggingface.co/papers/2307.14460).
        rA   r   r   r   ij)indexing)	start_dimN)r_   dtyper\   )r   r   )
rG   rD   stackmeshgridarangerg   
contiguousrN   r   sum)r   r   window_areacoords_flattenrelative_coordsrelative_position_indexs         r)    generate_relative_position_indexz9BeitRelativePositionBias.generate_relative_position_index   s    "#[^!3a!7AA<NQR<R SVW W!!n{1~5 KKu||KN'CU\\R]^_R`Ealpqr
 *!Q*5q$PQz8RR[[\]_`bcdooq1a KNQ$66 1a KNQ$66 1a AA$6$:: "'++K!O3E3IQ`QfQf"g*9*=*=b*AAB')>)B12&)>)BA&(=(A%&&r(   rr   c                    d| j                   d   z  dz
  }d| j                   d   z  dz
  }d|d   z  dz
  }d|d   z  dz
  }| j                  }| j                  }	||z  dz   }
|d|	dz
   }|j                  d||d      j	                  dddd      }t
        j                  j                  |t        |      t        |      fd      }|j	                  dddd      j                  |
dz
  d      }t        j                  |||	dz
  d g      }| j                  |      }||j                  d         }|j                  |d   |d   z  dz   |d   |d   z  dz   d      }|j	                  ddd      j                         }|rCt
        j                  j                  |j                  d      ||fdd	
      j                  d      }|j                  d      S )zu
        Modification of timm.models.beit.py: Attention._get_rel_pos_bias to support arbitrary window sizes.
        rA   r   r   r   Nr\   bilinear)r_   r`   Fr^   )r   r   r   rf   rg   r   rh   ri   r   rG   rk   r   rj   r   rw   squeeze)r;   r   rr   dim_size
old_height	old_widthro   rp    old_relative_position_bias_tableold_num_relative_distancenew_num_relative_distanceold_sub_tablenew_sub_table new_relative_position_bias_tabler   relative_position_biass                   r)   rF   z BeitRelativePositionBias.forward   s-    ))!,,q0
((++a/	Q'!+
A&*	+/+L+L($($>$>!$.$:Q$>!89X;TWX;XY%--aJKSSTUWXZ[]^_11:!6	)8L MT^ 2 
 &--aAq9AAB[^_B_acd+099<=VYZ=Z=\]^,
( #'"G"G"T!ABYB^B^_aBb!c "8!<!<N[^+a/Q+a.1PST1TVX"
 "8!?!?1a!H!S!S!U#%']]%>%>&003)#	 &? &
 gaj # &//22r(   )FN)r#   r$   r%   r   r2   staticmethodr   r   r   rG   r   r   boolrF   rH   rI   s   @r)   r   r      sk    	
z 	
d 	
 (4'eCHo '%,, ' 5 '4-3T -3]b]i]i -3r(   r   modulequerykeyvalueattention_maskscalingrW   kwargsc                    ||j                  d      dz  }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\         rA   r   )rc   r   )ptrainingr   )r_   rG   matmulrE   r   rh   softmaxfloat32tor   rW   r   r   )
r   r   r   r   r   r   rW   r   attn_weightsattn_outputs
             r)   eager_attention_forwardr     s     **R.D( <<s}}Q':;gEL!#n4==((2U]](SVVW\WbWbcL==((6??([L,,|U3K''1-88:K$$r(   c                        e Zd Zdef fdZ	 d	dej                  dej                  dz  dee   de	ej                  ej                  f   fdZ
 xZS )
BeitAttentionr,   c                    t         |           || _        |j                  | _        t	        |d|j
                  |j                  z        | _        |j                  | _        | j                  dz  | _	        d| _
        t        j                  |j
                  |j                  | j                  z        | _        t        j                  |j
                  |j                  | j                  z  d      | _        t        j                  |j
                  |j                  | j                  z        | _        t        j                  |j                  | j                  z  |j
                        | _        y )Nhead_dimr   F)bias)r1   r2   r,   r   getattrr9   r   attention_probs_dropout_probattention_dropoutr   	is_causalr   Linearq_projk_projv_projo_projr;   r,   r<   s     r)   r2   zBeitAttention.__init__)  s   #)#=#= 
F4F4F&JdJd4de!'!D!D}}d*ii 2 2F4N4NQUQ^Q^4^_ii 2 2F4N4NQUQ^Q^4^ejkii 2 2F4N4NQUQ^Q^4^_ii : :T]] JFL^L^_r(   Nhidden_statesr   r   r>   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                  sdn| j                  | j                  d|\  }
} |
j                  g |d j!                         }
| j#                  |
      }
|
|fS )Nr\   r   rA           )rW   r   )rB   r   r   rj   rE   r   r   r   get_interfacer,   _attn_implementationr   r   r   r   rf   r   r   )r;   r   r   r   input_shapehidden_shapequery_states
key_statesvalue_statesattention_interfacer   r   s               r)   rF   zBeitAttention.forward6  sK    $))#2.88b8$--8{{=166|DNNqRST[[/44\BLLQPQR
{{=166|DNNqRST(?(M(MKK,,.E)
 %8	%
  $}}C$2H2HLL	%
 	%
!\ *k));;;;FFHkk+.L((r(   r   )r#   r$   r%   r   r2   rG   r   r   r   r   rF   rH   rI   s   @r)   r   r   (  sf    `z `  /3)||) t+) +,	)
 
u||U\\)	*)r(   r   c                   \     e Zd Zdef fdZdej                  dej                  fdZ xZS )BeitMLPr,   c                    t         |           || _        t        |j                     | _        t        j                  |j                  |j                        | _
        t        j                  |j                  |j                        | _        y r   )r1   r2   r,   r
   
hidden_actactivation_fnr   r   r9   intermediate_sizefc1fc2r   s     r)   r2   zBeitMLP.__init__Y  sd    #F$5$5699V//1I1IJ99V55v7I7IJr(   r   r>   c                 l    | j                  |      }| j                  |      }| j                  |      }|S r   )r   r   r   r;   r   s     r)   rF   zBeitMLP.forward`  s4    /**=9/r(   	r#   r$   r%   r   r2   rG   r   rF   rH   rI   s   @r)   r   r   X  s,    Kz KU\\ ell r(   r   c                   r     e Zd ZdZd	deddf fdZdej                  dej                  fdZde	fdZ
 xZS )
BeitDropPathzStochastic depth (DropPath) per sample, for residual blocks.

    Identity when ``drop_prob`` is 0 or outside training. See `Deep Networks with Stochastic Depth
    <https://arxiv.org/abs/1603.09382>`_.
    	drop_probr>   Nc                 0    t         |           || _        y r   )r1   r2   r   )r;   r   r<   s     r)   r2   zBeitDropPath.__init__o  s    "r(   r   c                 P   | j                   dk(  s| j                  s|S d| j                   z
  }|j                  d   fd|j                  dz
  z  z   }t	        j
                  ||j                  |j                        }t	        j                  ||z         }|j                  |      |z  S )Nr   r   r   )r   )r   device)
r   r   rB   ndimrG   randr   r   floordiv)r;   r   	keep_probrB   random_tensors        r)   rF   zBeitDropPath.forwards  s    >>S   &	$$Q')DM4F4F4J,KK

50C0CML`L`aMI$=>  +m;;r(   c                      d| j                    S )Nzp=)r   r;   s    r)   
extra_reprzBeitDropPath.extra_repr|  s    DNN#$$r(   r   )r#   r$   r%   r&   floatr2   rG   r   rF   strr   rH   rI   s   @r)   r   r   h  sB    #% #$ #<U\\ <ell <%C %r(   r   c                        e Zd ZdZddedef fdZ	 	 	 ddej                  dej                  dz  de	d	e
eef   dz  d
ee   dej                  fdZ xZS )	BeitLayerz?This corresponds to the Block class in the timm implementation.r,   drop_path_ratec                 J   t         |           t        |      | _        t	        j
                  |j                  |j                        | _        t	        j
                  |j                  |j                        | _	        t        |      | _        t	        j                  |j                        | _        |j                  | _        |dkD  rt!        |      nt	        j"                         | _        |j&                  }|dkD  r7t	        j(                  |t+        j,                  |j                        z  d      nd| _        |dkD  r7t	        j(                  |t+        j,                  |j                        z  d      nd| _        |j2                  rt5        |      | _        y d | _        y )Nepsr   r   T)requires_gradg      ?)r1   r2   r   	attentionr   	LayerNormr9   layer_norm_epslayernorm_beforelayernorm_afterr   mlprU   rV   rW   r4   r   Identity	drop_pathlayer_scale_init_valuerM   rG   oneslambda_1lambda_2use_relative_position_biasr   r   )r;   r,   r   init_valuesr<   s       r)   r2   zBeitLayer.__init__  s@   &v. "V-?-?VEZEZ [!||F,>,>FDYDYZ6?zz&"<"<= ++9G#9Mn5SUS^S^S`33^ilm^mBLLuzz&2D2D'EEUYZsv 	 _jlm^mBLLuzz&2D2D'EEUYZsv 	 KQJkJk&>v&F#qu#r(   Nr   r   rr   
resolutionr   r>   c                 &   | j                   M|\  }}|| j                  z  || j                  z  f}| j                  |||j                  d         }	||	|z   n|	}|}
| j                  |      } | j                  |fd|i|\  }}| j                  |      }| j                  |z  }| j                  |      |
z   }|}
| j                  |      }| j                  |      }| j                  |      }| j                  |z  }| j                  |      |
z   }|S )Nr   )r   r   )r   r4   rB   r   r   rW   r  r  r   r  r  )r;   r   r   rr   r
  r   rY   rZ   r   r   residualry   s               r)   rF   zBeitLayer.forward  sE    &&2&MFE!T__4et6NOK%)%@%@5@S@STU@V &A &" <J;U&7[q 
 !--m<)4>>
)
 
q
 ]35}5@ !,,];/]35}5@r(   r   NFN)r#   r$   r%   r&   r   r   r2   rG   r   r   r   r   r   r   rF   rH   rI   s   @r)   r   r     s    Ivz v5 v, /3).-1&||& t+& #'	&
 #s(Od*& +,& 
&r(   r   c                        e Zd ZU eed<   dZdZdZdZdgZ	dZ
dZdZdZdZeedZd	Zd
gZ ej*                          fd       Z xZS )BeitPreTrainedModelr,   beitr=   )imageTr   F)r   
attentionsrR   z.*relative_position_index.*c                    t         |   |       t        |t              rwt	        j
                  |j                         |j                  t	        j
                  |j                         |j                   t	        j
                  |j                         yyt        |t              r t	        j
                  |j                         yt        |t              rt        |j                  t        j                        rit	        j                  |j                  | j                   j"                         t	        j                  |j$                  | j                   j"                         yyy)zInitialize the weightsN)r1   _init_weightsr5   rK   initzeros_rO   rQ   rT   r   r   r   r  r   rM   	constant_r,   r  r  )r;   r   r<   s     r)   r  z!BeitPreTrainedModel._init_weights  s     	f%fn-KK(()  ,F--.))5F667 6 89KK;;<	*&//2<<8v0R0RSv0R0RS 9 +r(   )r#   r$   r%   r   __annotations__base_model_prefixmain_input_nameinput_modalitiessupports_gradient_checkpointing_no_split_modules_supports_sdpa_supports_flash_attn_supports_flex_attn_supports_attention_backend_can_compile_fullgraphr   r   _can_record_outputs_input_embed_layer"_keys_to_ignore_on_load_unexpectedrG   no_gradr  rH   rI   s   @r)   r  r    s    $O!&*#$N "&!"# ,*H)I&U]]_T Tr(   r  c                        e Zd Zddededdf fdZe ed      e	 	 	 dde	j                  d	e	j                  dz  d
ede	j                  dz  dee   defd                     Z xZS )	BeitModelr,   add_pooling_layerr>   Nc           	         t         |   |       || _        t        |      | _        |j
                  rt        |      nd| _        t        |j                        D cg c]+  }|j                  |z  t        |j                  dz
  d      z  - }}t        j                  |D cg c]  }t        ||       c}      | _        |j                   rt        j"                         n*t        j$                  |j&                  |j(                        | _        |rt-        |      nd| _        | j1                          yc c}w c c}w )zv
        add_pooling_layer (bool, *optional*, defaults to `True`):
            Whether to add a pooling layer
        Nr   )r   r   )r1   r2   r,   rK   rX   !use_shared_relative_position_biasr   shared_position_biasrangenum_hidden_layersr   maxr   
ModuleListr   layersuse_mean_poolingr  r   r9   r   	layernorm
BeitPoolerpooler	post_init)r;   r,   r)  idrop_path_ratesrr<   s         r)   r2   zBeitModel.__init__  s   
 	 (0060X0X$V,^b 	! W\\b\t\tVu
QRF!!A%F,D,Dq,H!(LL
 
 mmRa$bQYva%H$bc $44BKKM",,vGYGY_e_t_t:u 	 ->j(4 	
 %cs   0D7"D<F)tie_last_hidden_statesr=   rs   rr   r   r   c                 
   | j                  ||      }|j                  dd }t        | j                  ||      }| j                  a|\  }}	|| j                  j
                  z  |	| j                  j
                  z  f}
| j	                  |
||j                  d         }|||z   n|}|}| j                  D ]  } ||f|||d|} | j                  |      }| j                  | j                  |      nd}t        ||      S )	z
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, num_patches)`, *optional*):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).
        )rs   rA   N)r,   inputs_embedsr   r   )rr   r   )r   rr   r
  )last_hidden_statepooler_output)
rX   rB   r   r,   r,  r4   r1  r3  r5  r"   )r;   r=   rs   rr   r   r   embedding_outputr
  rY   rZ   r   shared_relative_position_biasr   layersequence_outputpooled_outputs                   r)   rF   zBeitModel.forward   s<     ??<?Y!''+
2;;*)
 $$0&MFE!T[[%;%;;UdkkF\F\=\]K,0,E,E6NYiYoYopqYr -F -)
 "- .>2  )[[ 	E!-)A%	
 M	 ..78<8OO4UY)O[hiir(   )Tr  )r#   r$   r%   r   r   r2   r   r   r   rG   r   r   r   r   r"   rF   rH   rI   s   @r)   r(  r(    s    z d d 2  E2 48)..2-jll-j ))D0-j #'	-j
 t+-j +,-j 
$-j  3  -jr(   r(  c                   `     e Zd Zdeddf fdZdej                  dej                  fdZ xZS )r4  r,   r>   Nc                     t         |           |j                  r1t        j                  |j
                  |j                        | _        y d | _        y )Nr   )r1   r2   r2  r   r   r9   r   r3  r   s     r)   r2   zBeitPooler.__init__4  sA    KQKbKbBLL++1F1FG 	hl 	r(   r   c                     | j                   ,| j                  |d d dd d d f   j                  d            S |d d df   S )Nr   r   )r3  meanr   s     r)   rF   zBeitPooler.forward:  sD    BF..B\t~~mAqr1H5::1=>ubopqstptbuur(   r   rI   s   @r)   r4  r4  3  s4    
z 
d 
vU\\ vell vr(   r4  a  
    Beit Model transformer with a 'language' modeling head on top. BEiT does masked image modeling by predicting
    visual tokens of a Vector-Quantize Variational Autoencoder (VQ-VAE), whereas other vision models like ViT and DeiT
    predict RGB pixel values. As a result, this class is incompatible with [`AutoModelForMaskedImageModeling`], so you
    will need to use [`BeitForMaskedImageModeling`] directly if you wish to do masked image modeling with BEiT.
    c                        e Zd Zdeddf fdZd Zee	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  d	ed
ej                  dz  dee   deez  fd              Z xZS )BeitForMaskedImageModelingr,   r>   Nc                 H   t         |   |       |j                  | _        t        |d      | _        t        j                  |j                  |j                        | _	        t        j                  |j                  |j                        | _        | j                          y )NFr)  r   )r1   r2   
num_labelsr(  r  r   r   r9   r   r3  r   
vocab_sizelm_headr6  r   s     r)   r2   z#BeitForMaskedImageModeling.__init__H  su      ++f>	 f&8&8f>S>STyy!3!3V5F5FG 	r(   c                      y r   r'   r   s    r)   get_output_embeddingsz0BeitForMaskedImageModeling.get_output_embeddingsU  s    r(   r=   rs   labelsrr   r   r   c                 ,    | j                   |f|||d|}|j                  }| j                  |      }| j                  |ddddf         }	d}
| t	        j
                         } ||	|   |      }
t        |
|	|j                  |j                        S )a  
        bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, num_patches)`):
            Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the image classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).

        Examples:

        ```python
        >>> from transformers import AutoImageProcessor, BeitForMaskedImageModeling
        >>> import torch
        >>> from PIL import Image
        >>> import requests

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = Image.open(requests.get(url, stream=True).raw)

        >>> image_processor = AutoImageProcessor.from_pretrained("microsoft/beit-base-patch16-224-pt22k")
        >>> model = BeitForMaskedImageModeling.from_pretrained("microsoft/beit-base-patch16-224-pt22k")

        >>> num_patches = (model.config.image_size // model.config.patch_size) ** 2
        >>> pixel_values = image_processor(images=image, return_tensors="pt").pixel_values
        >>> # create random boolean mask of shape (batch_size, num_patches)
        >>> bool_masked_pos = torch.randint(low=0, high=2, size=(1, num_patches)).bool()

        >>> outputs = model(pixel_values, bool_masked_pos=bool_masked_pos)
        >>> loss, logits = outputs.loss, outputs.logits
        >>> list(logits.shape)
        [1, 196, 8192]
        ```)rs   rr   r   Nr   losslogitsr   r  )	r  r=  r3  rN  r   CrossEntropyLossr   r   r  )r;   r=   rs   rQ  rr   r   r   outputsrB  prediction_scoresmasked_lm_lossloss_fcts               r)   rF   z"BeitForMaskedImageModeling.forwardX  s    X $))
+%=)	

 
 "33..9 LLAB)?@**,H%&7&H&QN$!//))	
 	
r(   )NNNFN)r#   r$   r%   r   r2   rP  r   r   rG   r   r   r   r   r   r   r   rF   rH   rI   s   @r)   rI  rI  ?  s    z d   -137&*)..2@
llT)@
 ))D0@
 t#	@

 #'@
 t+@
 +,@
 
	@
  @
r(   rI  z
    Beit Model transformer with an image classification head on top (a linear layer on top of the average of the final
    hidden states of the patch tokens) e.g. for ImageNet.
    c                        e Zd Zdeddf fdZee	 	 	 d
dej                  dz  dej                  dz  de	de
e   deez  f
d	              Z xZS )BeitForImageClassificationr,   r>   Nc                 .   t         |   |       |j                  | _        t        |d      | _        |j                  dkD  r*t        j                  |j                  |j                        nt        j                         | _	        | j                          y )NTrK  r   )r1   r2   rL  r(  r  r   r   r9   r  
classifierr6  r   s     r)   r2   z#BeitForImageClassification.__init__  ss      ++f=	 OUN_N_bcNc"))F$6$68I8IJikititiv 	r(   r=   rQ  rr   r   c                      | j                   |fd|i|}|j                  }| j                  |      }d}|| j                  ||| j                        }t        |||j                  |j                        S )a  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the image classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
        rr   NrS  )r  r>  r^  loss_functionr,   r   r   r  )	r;   r=   rQ  rr   r   rW  rC  rU  rT  s	            r)   rF   z"BeitForImageClassification.forward  s     $))
%=
 
  --/%%ffdkkBD$!//))	
 	
r(   NNF)r#   r$   r%   r   r2   r   r   rG   r   r   r   r   r   r   rF   rH   rI   s   @r)   r\  r\    s    
z 
d 
  -1&*).	 
llT) 
 t# 
 #'	 

 +, 
 
&	& 
   
r(   r\  c                        e Zd Z	 	 	 	 	 	 	 ddededeeeef   z  dedeeeef   z  ez  dedeeeef   z  ded	ef fd
Zdej                  dej                  fdZ
 xZS )BeitConvLayerin_channelsout_channelsr/   r0   paddingr   dilationgroups
activationc
           
          t         
|           t        j                  ||||||||      | _        t        j
                  |      | _        |	t        |	   | _	        y t        j                         | _	        y )N)rd  re  r/   r0   rf  rg  rh  r   )
r1   r2   r   r8   convolutionBatchNorm2dnormalizationr
   r  ri  )r;   rd  re  r/   r0   rf  r   rg  rh  ri  r<   s             r)   r2   zBeitConvLayer.__init__  si     	99#%#	
  ^^L90:0F&,BKKMr(   r   r>   c                 l    | j                  |      }| j                  |      }| j                  |      }|S r   )rk  rm  ri  r   s     r)   rF   zBeitConvLayer.forward  s6    ((7**=96r(   )r   r   r   Fr   r   relu)r#   r$   r%   r   r   r   r   r2   rG   r   rF   rH   rI   s   @r)   rc  rc    s    
 .//0*+ ZZ Z 5c?*	Z
 Z uS#X&,Z Z c3h'Z Z Z4U\\ ell r(   rc  c                   v     e Zd Zdedededdf fdZdej                  deeef   dej                  fd	Z xZ	S )
BeitPyramidPoolingBlock
pool_scalerd  channelsr>   Nc                 |    t         |           t        j                  |      | _        t        ||d      | _        y )Nr   r/   )r1   r2   r   AdaptiveAvgPool2dpoolingrc  conv)r;   rr  rd  rs  r<   s       r)   r2   z BeitPyramidPoolingBlock.__init__  s0    ++J7!+xQG	r(   inputr_   c                     | j                  |      }| j                  |      }t        j                  j	                  ||dd      }|S )Nr   Fr^   )rw  rx  r   rh   ri   )r;   ry  r_   hidden_states       r)   rF   zBeitPyramidPoolingBlock.forward  sB    ||E*yy.}}00Dzin0or(   )
r#   r$   r%   r   r2   rG   r   r   rF   rH   rI   s   @r)   rq  rq    sS    H3 HS HC HD H
U\\ sCx U\\ r(   rq  c                   |     e Zd ZdZdeedf   dededdf fdZd	ej                  de	ej                     fd
Z
 xZS )BeitPyramidPoolingModuleak  
    Pyramid Pooling Module (PPM) used in PSPNet.

    Args:
        pool_scales (tuple[int]): Pooling scales used in Pooling Pyramid
            Module.
        in_channels (int): Input channels.
        channels (int): Channels after modules, before conv_seg.

    Based on OpenMMLab's implementation, found in https://github.com/open-mmlab/mmsegmentation.
    pool_scales.rd  rs  r>   Nc           
          t         |           || _        || _        || _        t        j                  |D cg c]  }t        |||       c}      | _        y c c}w )N)rr  rd  rs  )	r1   r2   r~  rd  rs  r   r0  rq  blocks)r;   r~  rd  rs  rr  r<   s        r)   r2   z!BeitPyramidPoolingModule.__init__  s\    && mm #. (:;aij
s   Ar   c                 v    |j                         dd  }| j                  D cg c]  } |||       c}S c c}w )NrA   )r_   )r_   r  )r;   r   original_sizeblocks       r)   rF   z BeitPyramidPoolingModule.forward  s6    %**,QR0FJkkRUm-8RRRs   6)r#   r$   r%   r&   r   r   r2   rG   r   r   rF   rH   rI   s   @r)   r}  r}    sV    


E#s(O 

# 

QT 

Y] 

SU\\ Sd5<<6H Sr(   r}  c                        e Zd ZdZdeddf fdZdeej                     dej                  fdZ	deej                     dej                  fd	Z
 xZS )
BeitUperHeadz
    Unified Perceptual Parsing for Scene Understanding. This head is the implementation of
    [UPerNet](https://huggingface.co/papers/1807.10221).

    Based on OpenMMLab's implementation, found in https://github.com/open-mmlab/mmsegmentation.
    r,   r>   Nc           	         t         |           |j                  | _        |j                  gdz  | _        |j                  | _        t        j                  | j
                  |j                  d      | _	        t        | j                  | j                  d   | j
                        | _        t        | j                  d   t        | j                        | j
                  z  z   | j
                  dd      | _        t        j                         | _        t        j                         | _        | j                  d d D ]o  }| j                   j%                  t        || j
                  d             | j"                  j%                  t        | j
                  | j
                  dd             q t        t        | j                        | j
                  z  | j
                  dd      | _        y )N   r   ru  r\   r   r/   rf  )r1   r2   r~  r9   rd  rs  r   r8   rL  r^  r}  psp_modulesrc  lenpsp_bottleneckr0  lateral_convs	fpn_convsappendfpn_bottleneck)r;   r,   rd  r<   s      r)   r2   zBeitUperHead.__init__*  s   !--"../!3**))DMM63D3DRST 4R MM

 ,R 3t'7'7#84==#HHMM	
  ]]_++CR0 	iK%%mK\]&^_NN!!-t}}Z[ef"gh	i ,  !DMM1MM	
r(   r   c                     |d   }t        j                  |g| j                  |      d      }| j                  |      S ru   )rG   rk   r  r  )r;   r   r{  s      r)   psp_forwardzBeitUperHead.psp_forwardL  sA    $R(yy,!P1A1A,1O!PVWX""<00r(   encoder_hidden_statesc                 6   g }t        | j                  |      D ]  \  }}|j                   ||              |j                  | j                  |             t	        |      }t        |dz
  dd      D ]L  }||dz
     j                  dd  }||dz
     t        j                  j                  ||   |dd      z   ||dz
  <   N g }t        |dz
        D ])  }|j                   | j                  |   ||                + |j                  |d          t        |dz
  dd      D ];  }t        j                  j                  ||   |d   j                  dd  dd      ||<   = t        j                  |d      }| j                  |      }	| j                  |	      }	|	S )	Nr   r   r\   rA   r   Fr^   rb   )zipr  r  r  r  r-  rB   r   rh   ri   r  rG   rk   r  r^  )
r;   r  lateralslateral_convr{  used_backbone_levelsr7  
prev_shapefpn_outsoutputs
             r)   rF   zBeitUperHead.forwardQ  s   *-d.@.@BW*X 	8&L,OOL67	8 	(()>?@  #8}+a/B7 	A!!a%..qr2J&q1uo0I0I*:U 1J 1 HQUO	 +a/0 	<AOO-DNN1-hqk:;	< 	%+a/B7 	A--33(1+"3"3AB"7jX] 4 HQK	 99X1-$$X.(r(   )r#   r$   r%   r&   r   r2   r   rG   r   r  rF   rH   rI   s   @r)   r  r  "  s\     
z  
d  
D1ell); 1 1
T%,,-? ELL r(   r  c                        e Zd ZdZ	 ddedededeeeef   z  ddf
 fdZd	ee	j                     de	j                  fd
Z xZS )BeitFCNHeada  
    Fully Convolution Networks for Semantic Segmentation. This head is implemented of
    [FCNNet](https://huggingface.co/papers/1411.4038>).

    Args:
        config (BeitConfig): Configuration.
        in_channels
        kernel_size (int): The kernel size for convs in the head. Default: 3.
        dilation (int): The dilation rate for convs in the head. Default: 1.


    Based on OpenMMLab's implementation, found in https://github.com/open-mmlab/mmsegmentation.
    r,   in_indexr/   rg  r>   Nc           
      0   t         |           |j                  | _        |j                  | _        |j                  | _        |j                  | _	        || _
        |dz  |z  }t        j                         | _        | j                  dkD  r| j                  j                  t        | j                  | j
                  |||             t!        | j                  dz
        D ]?  }| j                  j                  t        | j
                  | j
                  |||             A | j                  r8t        | j                  | j
                  z   | j
                  ||dz        | _        t        j$                  | j
                  |j&                  d      | _        y )NrA   r   )r/   rf  rg  r   r  ru  )r1   r2   r9   rd  auxiliary_channelsrs  auxiliary_num_convs	num_convsauxiliary_concat_inputconcat_inputr  r   r0  convsr  rc  r-  conv_catr8   rL  r^  )r;   r,   r  r/   rg  conv_paddingry   r<   s          r)   r2   zBeitFCNHead.__init__  sQ    	!--1133"99 #q(H4]]_
>>AJJ$$dmmVbmu
 4>>A-. 	

!!!$/ ,!)	 )  4==0$--[bmqrbrDM ))DMM63D3DRSTr(   r  c                     || j                      }|}| j                  D ]
  } ||      } | j                  r(| j                  t	        j
                  ||gd            }| j                  |      }|S )Nr   rb   )r  r  r  r  rG   rk   r^  )r;   r  r  r   rx  s        r)   rF   zBeitFCNHead.forward  sn    (7 JJ 	0D /M	0 MM%))X}4MST*UVM6r(   )rA   r   r   )r#   r$   r%   r&   r   r   r   r2   r   rG   r   rF   rH   rI   s   @r)   r  r  s  su     no!U !U,/!UBE!UUX[`adfiai[jUj!U	!UFT%,,-? ELL r(   r  c            	       n     e Zd ZdZd
dedededdf fdZdej                  dej                  fd	Z xZ	S )BeitFPNUpBlockuE   4x upsampling block: ConvTranspose → BN → GELU → ConvTranspose.r9   r/   r0   r>   Nc                     t         |           t        j                  ||||      | _        t        j
                  |      | _        t        j                         | _        t        j                  ||||      | _	        y )Nr.   )
r1   r2   r   ConvTranspose2dconv_transpose1rl  rm  GELUri  conv_transpose2)r;   r9   r/   r0   r<   s       r)   r2   zBeitFPNUpBlock.__init__  sb    !11+{Xclrs^^K8'')!11+{Xclrsr(   r   c                     | j                  |      }| j                  |      }| j                  |      }| j                  |      }|S r   )r  rm  ri  r  r   s     r)   rF   zBeitFPNUpBlock.forward  sF    ,,];**=96,,];r(   )rA   rA   )
r#   r$   r%   r&   r   r2   rG   r   rF   rH   rI   s   @r)   r  r    sH    OtC tc ts tSW tU\\ ell r(   r  c                   t     e Zd ZdZdef fdZdeej                  df   deej                  df   fdZ	 xZ
S )BeitFPNNeckz
    4-level feature pyramid neck for BeiT. Produces x4 upsample, x2 upsample,
    identity, and x2 downsample outputs from the four selected ViT feature maps.
    r,   c                     t         |           t        |j                        | _        t        j                  |j                  |j                  dd      | _        t        j                  dd      | _	        y )NrA   r.   )
r1   r2   r  r9   fpn1r   r  fpn2	MaxPool2dfpn4r   s     r)   r2   zBeitFPNNeck.__init__  sX    "6#5#56	&&v'9'96;M;M[\efg	LLQq9	r(   feature_maps.r>   c                     | j                  |d         | j                  |d         |d   | j                  |d         fS r   )r  r  r  )r;   r  s     r)   rF   zBeitFPNNeck.forward  sC    IIl1o&IIl1o&OIIl1o&	
 	
r(   )r#   r$   r%   r&   r   r2   r   rG   r   rF   rH   rI   s   @r)   r  r    sD    
:z :
E%,,*;$< 
u||UXGXAY 
r(   r  c                        e Zd Zdeddf fdZeee	 	 	 d
dej                  dz  dej                  dz  de
dee   deez  f
d	                     Z xZS )BeitForSemanticSegmentationr,   r>   Nc                 `   t         |   |       |j                  | _        t        |d      | _        t        | j                  j                        dk7  rt        d      t        |      | _
        t        |      | _        |j                  rt        |      nd | _        | j!                          y )NFrK  r  zBeitForSemanticSegmentation requires config.out_indices to be a list of 4 integers, specifying which features to use from the backbone. One can use [3, 5, 7, 11] in case of a base-sized architecture.)r1   r2   rL  r(  r  r  r,   out_indicesrC   r  fpnr  decode_headuse_auxiliary_headr  auxiliary_headr6  r   s     r)   r2   z$BeitForSemanticSegmentation.__init__  s      ++f>	t{{&&'1,- 
 v& (/5;5N5Nk&1TX 	r(   r=   rQ  rr   r   c                    |$| j                   j                  dk(  rt        d       | j                  |fd|i|}|j                  |j
                  \  }}}|| j                   j                  z  || j                   j                  z  t        fd| j                   j                  D              }	| j                  |	      }	| j                  |	      }
d}| j                  | j                  |	      }d}|>| j                  |
|| j                   j                  || j                   j                        }t        ||
|j                  |j                         S )aD  
        labels (`torch.LongTensor` of shape `(batch_size, height, width)`, *optional*):
            Ground truth semantic segmentation maps for computing the loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels > 1`, a classification loss is computed (Cross-Entropy).

        Examples:

        ```python
        >>> from transformers import AutoImageProcessor, BeitForSemanticSegmentation
        >>> from PIL import Image
        >>> import requests

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = Image.open(requests.get(url, stream=True).raw)

        >>> image_processor = AutoImageProcessor.from_pretrained("microsoft/beit-base-finetuned-ade-640-640")
        >>> model = BeitForSemanticSegmentation.from_pretrained("microsoft/beit-base-finetuned-ade-640-640")

        >>> inputs = image_processor(images=image, return_tensors="pt")
        >>> outputs = model(**inputs)
        >>> # logits are of shape (batch_size, num_labels, height, width)
        >>> logits = outputs.logits
        ```Nr   z/The number of labels should be greater than onerr   c              3      K   | ]7  }|d z
     ddd df   j                  d d      j                  d       9 yw)r   NrA   r\   )rE   rf   ).0r7  rz   r  patch_heightpatch_widths     r)   	<genexpr>z6BeitForSemanticSegmentation.forward.<locals>.<genexpr>  sN      
 "!a%(AB/99!Q?GG
TVXdfqr
s   =A )ignore_indexauxiliary_logitsauxiliary_loss_weightrS  )r,   rL  rC   r  r   rB   r4   r   r  r  r  r  r`  semantic_loss_ignore_indexr  r   r  )r;   r=   rQ  rr   r   rW  ry   rY   rZ   r  rU  r  rT  rz   r  r  r  s                @@@@r)   rF   z#BeitForSemanticSegmentation.forward  sh   B $++"8"8A"=NOO$))
%=
 
 !( 5 5'3'9'9$
Avu!7!77t{{555  
[[,,
 
 xx-!!,/*#22<@%%![[CC!1&*kk&G&G & D '!//))	
 	
r(   ra  )r#   r$   r%   r   r2   r   r   r   rG   r   r   r   r   r   r   rF   rH   rI   s   @r)   r  r    s    z d *   -1&*).	G
llT)G
 t#G
 #'	G

 +,G
 
(	(G
  ! G
r(   r  zM
    BEiT backbone, to be used with frameworks like DETR and MaskFormer.
    c            	       V     e Zd Z fdZeeededee	   de
fd                     Z xZS )BeitBackbonec                 <   t         |   |       t        |j                  dz         D cg c]  }|j                   c}| _        t        |d      | _        |j                  rt        |      nt        j                         | _        | j                          y c c}w )Nr   FrK  )r1   r2   r-  r.  r9   num_featuresr(  r  add_fpnr  r   r  r  r6  )r;   r,   ry   r<   s      r)   r2   zBeitBackbone.__init__A  st     9>v?W?WZ[?[9\]AV//]f>	*0..;v&bkkm 	 ^s   Br=   r   r>   c                 *   |j                   \  }}}}|| j                  j                  z  }|| j                  j                  z  } | j                  |fi |}	|	j                  }
d}t        | j                  |
      D ]d  \  }}|| j                  v s| j                  j                  r4|ddddddf   }|j                  dd      }|j                  |d||      }||fz  }f | j                  |      }t        ||	j                  |	j                        S )a:  
        Examples:

        ```python
        >>> from transformers import AutoImageProcessor, AutoBackbone
        >>> import torch
        >>> from PIL import Image
        >>> import requests

        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = Image.open(requests.get(url, stream=True).raw)

        >>> processor = AutoImageProcessor.from_pretrained("microsoft/beit-base-patch16-224")
        >>> model = AutoBackbone.from_pretrained(
        ...     "microsoft/beit-base-patch16-224", out_features=["stage1", "stage2", "stage3", "stage4"]
        ... )

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

        >>> outputs = model(**inputs)
        >>> feature_maps = outputs.feature_maps
        >>> list(feature_maps[-1].shape)
        [1, 768, 14, 14]
        ```r'   Nr   rA   r\   )r  r   r  )rB   r,   r4   r  r   r  stage_namesout_featuresreshape_hidden_statesrE   rf   r  r   r  )r;   r=   r   rz   ry   rY   rZ   r  r  rW  r   r  stager{  s                 r)   rF   zBeitBackbone.forwardK  s   @ (4'9'9$
Avu!7!77t{{555$))L3F3--#&t'7'7#G 	0E<)));;44#/12q#9L#/#9#9!Q#?L#/#7#7
BVa#bL/	0 xx-%!//))
 	
r(   )r#   r$   r%   r2   r   r   r   r   r   r   r   rF   rH   rI   s   @r)   r  r  ;  sN      3
3
 +,3
 
	3
  ! 3
r(   r  )r\  rI  r  r(  r  r  )Nr   )Fcollections.abcr   r   dataclassesr   rG   r   r    r	   r  activationsr
   backbone_utilsr   r   masking_utilsr   modeling_layersr   modeling_outputsr   r   r   r   r   modeling_utilsr   r   processing_utilsr   pytorch_utilsr   utilsr   r   r   utils.genericr   r   utils.output_capturingr   configuration_beitr   r"   Moduler+   rK   r   r   r   r   r   r   r   r  r(  r4  rI  r\  rc  rq  r}  r  r  r  r  r  r  __all__r'   r(   r)   <module>r     s  * / !   & ! H 6 9  G & @ B B I 5 * 
 !;  H")) H<SRYY SlV3ryy V3~ !%II%<<% 
% <<	%
 LL4'% T\% % '(%8-)BII -)`bii  %299 %0<* <~ "T/ "T "TJ Jj# Jj JjZ	v 	v S
!4 S
S
l /
!4 /
/
dBII D
bii 
Sryy S<N299 Nb:")) :zRYY $
")) 
* `
"5 `
 `
F 
A
="5 A

A
Hr(   