
    ^jȖ                        d 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
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 ddlmZmZmZ ddlmZm Z  ddl!m"Z" ddl#m$Z$ ddl%m&Z& ddl'm(Z(m)Z)m*Z*m+Z+m,Z,m-Z- ddl.m/Z/  ed      e G d de                    Z0 G d de,      Z1 G d de)      Z2 G d dejf                        Z4 G d  d!e(      Z5 G d" d#e+      Z6 G d$ d%e&      Z7 G d& d'e*      Z8e G d( d)e-             Z9e G d* d+e9             Z: G d, d-ejf                        Z; ed.       G d/ d0e9             Z< ed1       G d2 d3e9             Z= G d4 d5e$      Z> G d6 d7ejf                        Z? G d8 d9ejf                        Z@ G d: d;ejf                        ZA G d< d=ejf                        ZB G d> d?ejf                        ZC G d@ dAejf                        ZDe G dB dCe9             ZE edD       G dE dFe
e9             ZFg dGZGy)HzPyTorch BEiT model.    )	dataclassN)Tensornn   )initialization)BackboneMixinfilter_output_hidden_states)create_bidirectional_mask)BackboneOutputBaseModelOutputWithPoolingImageClassifierOutputMaskedLMOutputSemanticSegmenterOutput)PreTrainedModel)Unpack)#compile_compatible_method_lru_cache)TransformersKwargsauto_docstring	torch_int)can_return_tuplemerge_with_config_defaults)capture_outputs   )ResNetConvLayer)SwinDropPath)ViTAttentionViTEmbeddingsViTLayerViTMLPViTPatchEmbeddingsViTPreTrainedModel   )
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__     p/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/beit/modular_beit.pyr&   r&   +   s    r,   r&   c                       e Zd Zy)BeitPatchEmbeddingsNr'   r(   r)   r+   r,   r-   r/   r/   :       r,   r/   c                   v    e Zd ZdeddfdZ	 ddej                  dej                  dz  dej                  fdZy)	BeitEmbeddingsconfigreturnNc                    t         j                  j                  |        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"   )r   Module__init__	Parametertorchzeroshidden_size	cls_tokenuse_mask_token
mask_tokenr/   patch_embeddings
patch_sizenum_patches use_absolute_position_embeddingsposition_embeddingsDropouthidden_dropout_probdropout)selfr4   rB   s      r-   r8   zBeitEmbeddings.__init__?   s    
		4 ekk!Q8J8J&KLQWQfQf",,u{{1a9K9K'LMlp 3F ; ++++77 66 LLQa9K9KLM 	 
 zz&"<"<=r,   pixel_values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"   dim)shaper@   sizer?   expand	unsqueezetype_asr=   r:   catrD   interpolate_pos_encodingrG   )rH   rI   rJ   _heightwidth
embeddings
batch_sizeseq_lenmask_tokensmask
cls_tokenss               r-   forwardzBeitEmbeddings.forwardN   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#   r8   r:   r   
BoolTensorr`   r+   r,   r-   r3   r3   >   sN    >z >d >$ 48ll ))D0 
	r,   r3   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 )BeitRelativePositionBiasr4   r5   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"   r   r   )superr8   
image_size
isinstancetuplelistrA   window_sizenum_relative_distancer   r9   r:   r;   num_attention_headsrelative_position_bias_table)rH   r4   rh   	__class__s      r-   r8   z!BeitRelativePositionBias.__init__i   s    &&
*udm4$j1J&qMV->->>
1QWQbQb@bc&'$*:*:1*=&=&Aa$JZJZ[\J]F]`aFa%bef%f",.LLKK22F4N4NO-
)r,   
   )maxsizerl   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).
        r   r   r"   r   ij)indexing)	start_dimN)rQ   dtyperM   )r   r   )
r:   flattenstackmeshgridarangepermute
contiguousr;   rw   sum)rl   rm   window_areacoords_flattenrelative_coordsrelative_position_indexs         r-    generate_relative_position_indexz9BeitRelativePositionBias.generate_relative_position_indexu   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,   rV   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.
        r   r   r"   r   NrM   bilinear)rQ   modeFrQ   r   align_corners)rl   ro   rm   reshaper|   r   
functionalinterpolater   r:   rU   r   viewr}   rS   squeeze)rH   rl   rV   dim_size
old_height	old_width
new_height	new_width 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-   r`   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#   r8   staticmethodr   rj   intr:   r   r   boolr`   __classcell__rp   s   @r-   rd   rd   h   sk    	
z 	
d 	
 (4'eCHo '%,, ' 5 '4-3T -3]b]i]i -3r,   rd   c                   $     e Zd Zdef fdZ xZS )BeitAttentionr4   c                    t         |   |       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 )NF)bias)rg   r8   r   Linearr<   rn   head_dimq_projk_projv_projo_projrH   r4   rp   s     r-   r8   zBeitAttention.__init__   s     ii 2 2F4N4NQUQ^Q^4^_ii 2 2F4N4NQUQ^Q^4^ejkii 2 2F4N4NQUQ^Q^4^_ii : :T]] JFL^L^_r,   )r'   r(   r)   r#   r8   r   r   s   @r-   r   r      s    `z ` `r,   r   c                       e Zd Zy)BeitMLPNr0   r+   r,   r-   r   r      r1   r,   r   c                       e Zd Zy)BeitDropPathNr0   r+   r,   r-   r   r      r1   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.r4   drop_path_ratec                    t         |           |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| _        |j                  rt        |      | _        y d | _        y )N        r   T)requires_gradg      ?)rg   r8   rA   r   r   Identity	drop_pathlayer_scale_init_valuer9   r:   onesr<   lambda_1lambda_2use_relative_position_biasrd   r   )rH   r4   r   init_valuesrp   s       r-   r8   zBeitLayer.__init__   s     ++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,   Nhidden_statesattention_maskrV   
resolutionkwargsr5   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   rA   rP   layernorm_before	attentionrG   r   r   layernorm_aftermlpr   )rH   r   r   rV   r   r   rX   rY   rl   r   residualrW   s               r-   r`   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#   floatr8   r:   r   r   rj   r   r   r   r`   r   r   s   @r-   r   r      s    Ivz v5 v" /3).-1&||& t+& #'	&
 #s(Od*& +,& 
&r,   r   c                   &    e Zd ZdgZdgZdZdZd Zy)BeitPreTrainedModelr   z.*relative_position_index.*Fc                    t        j                  | |       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)r   _init_weightsri   r3   initzeros_r=   r?   rD   rd   ro   r   r   r   r9   	constant_r4   r   r   )rH   modules     r-   r   z!BeitPreTrainedModel._init_weights  s    %%dF3fn-KK(()  ,F--.))5F667 6 89KK;;<	*&//2<<8v0R0RSv0R0RS 9 +r,   N)r'   r(   r)   _no_split_modules"_keys_to_ignore_on_load_unexpected_supports_flash_attn_supports_flex_attnr   r+   r,   r-   r   r     s%    $*H)I& 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 )	BeitModelr4   add_pooling_layerr5   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   eps)rg   r8   r4   r3   rZ   !use_shared_relative_position_biasrd   shared_position_biasrangenum_hidden_layersr   maxr   
ModuleListr   layersuse_mean_poolingr   	LayerNormr<   layer_norm_eps	layernorm
BeitPoolerpooler	post_init)rH   r4   r   idrop_path_ratesrrp   s         r-   r8   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_statesrI   rJ   rV   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).
        )rJ   r   N)r4   inputs_embedsr   r"   )rV   r   )r   rV   r   )last_hidden_statepooler_output)
rZ   rP   r
   r4   r   rA   r   r   r   r&   )rH   rI   rJ   rV   r   r   embedding_outputr   rX   rY   rl   shared_relative_position_biasr   layersequence_outputpooled_outputs                   r-   r`   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   r8   r   r   r   r:   r   rb   r   r   r&   r`   r   r   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 )r   r4   r5   Nc                     t         |           |j                  r1t        j                  |j
                  |j                        | _        y d | _        y )Nr   )rg   r8   r   r   r   r<   r   r   r   s     r-   r8   zBeitPooler.__init__s  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   )r   meanrH   r   s     r-   r`   zBeitPooler.forwardy  sD    BF..B\t~~mAqr1H5::1=>ubopqstptbuur,   )	r'   r(   r)   r#   r8   r:   r   r`   r   r   s   @r-   r   r   r  s4    
z 
d 
vU\\ vell vr,   r   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 )BeitForMaskedImageModelingr4   r5   Nc                 H   t         |   |       |j                  | _        t        |d      | _        t        j                  |j                  |j                        | _	        t        j                  |j                  |j                        | _        | j                          y )NFr   r   )rg   r8   
num_labelsr   beitr   r   r<   r   r   r   
vocab_sizelm_headr   r   s     r-   r8   z#BeitForMaskedImageModeling.__init__  su      ++f>	 f&8&8f>S>STyy!3!3V5F5FG 	r,   c                      y ra   r+   )rH   s    r-   get_output_embeddingsz0BeitForMaskedImageModeling.get_output_embeddings  s    r,   rI   rJ   labelsrV   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]
        ```)rJ   rV   r   Nr"   losslogitsr   
attentions)	r   r   r   r  r   CrossEntropyLossr   r   r	  )rH   rI   rJ   r  rV   r   r   outputsr   prediction_scoresmasked_lm_lossloss_fcts               r-   r`   z"BeitForMaskedImageModeling.forward  s    X $))
+%=)	

 
 "33..9 LLAB)?@**,H%&7&H&QN$!//))	
 	
r,   )NNNFN)r'   r(   r)   r#   r8   r  r   r   r:   r   rb   r   r   r   rj   r   r`   r   r   s   @r-   r   r   ~  s    z d   -137&*)..2@
llT)@
 ))D0@
 t#	@

 #'@
 t+@
 +,@
 
	@
  @
r,   r   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 )BeitForImageClassificationr4   r5   Nc                 .   t         |   |       |j                  | _        t        |d      | _        |j                  dkD  r*t        j                  |j                  |j                        nt        j                         | _	        | j                          y )NTr   r   )rg   r8   r   r   r   r   r   r<   r   
classifierr   r   s     r-   r8   z#BeitForImageClassification.__init__  ss      ++f=	 OUN_N_bcNc"))F$6$68I8IJikititiv 	r,   rI   r  rV   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).
        rV   Nr  )r   r   r  loss_functionr4   r   r   r	  )	rH   rI   r  rV   r   r  r   r  r  s	            r-   r`   z"BeitForImageClassification.forward  s     $))
%=
 
  --/%%ffdkkBD$!//))	
 	
r,   NNF)r'   r(   r)   r#   r8   r   r   r:   r   r   r   r   rj   r   r`   r   r   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 xZS )BeitConvLayerin_channelsout_channelskernel_sizestridepaddingr   dilationgroups
activationc
           
      f    t         
|           t        j                  ||||||||      | _        y )N)r  r  r  r  r  r  r  r   )rg   r8   r   Conv2dconvolution)rH   r  r  r  r  r  r   r  r  r  rp   s             r-   r8   zBeitConvLayer.__init__  s9     	99#%#	
r,   )r   r"   r   Fr"   r"   relu)	r'   r(   r)   r   rj   strr   r8   r   r   s   @r-   r  r    s    
 .//0*+ 

 
 5c?*	

 
 uS#X&,
 
 c3h'
 
 
 
r,   r  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_scaler  channelsr5   Nc                 |    t         |           t        j                  |      | _        t        ||d      | _        y )Nr"   r  )rg   r8   r   AdaptiveAvgPool2dpoolingr  conv)rH   r'  r  r(  rp   s       r-   r8   z BeitPyramidPoolingBlock.__init__/  s0    ++J7!+xQG	r,   inputrQ   c                     | j                  |      }| j                  |      }t        j                  j	                  ||dd      }|S )Nr   Fr   )r,  r-  r   r   r   )rH   r.  rQ   hidden_states       r-   r`   zBeitPyramidPoolingBlock.forward4  sB    ||E*yy.}}00Dzin0or,   )
r'   r(   r)   r   r8   r:   r   rj   r`   r   r   s   @r-   r&  r&  .  sS    H3 HS HC HD H
U\\ sCx U\\ r,   r&  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.r  r(  r5   Nc           
          t         |           || _        || _        || _        t        j                  |D cg c]  }t        |||       c}      | _        y c c}w )N)r'  r  r(  )	rg   r8   r3  r  r(  r   r   r&  blocks)rH   r3  r  r(  r'  rp   s        r-   r8   z!BeitPyramidPoolingModule.__init__H  s\    && mm #. (:;aij
s   Ar   c                 v    |j                         dd  }| j                  D cg c]  } |||       c}S c c}w )Nr   )rQ   )rQ   r5  )rH   r   original_sizeblocks       r-   r`   z BeitPyramidPoolingModule.forwardT  s6    %**,QR0FJkkRUm-8RRRs   6)r'   r(   r)   r*   rj   r   r8   r:   r   rk   r`   r   r   s   @r-   r2  r2  ;  sV    


E#s(O 

# 

QT 

Y] 

SU\\ Sd5<<6H Sr,   r2  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.
    r4   r5   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"   r*  rM   r   r  r  )rg   r8   r3  r<   r  r(  r   r!  r   r  r2  psp_modulesr  lenpsp_bottleneckr   lateral_convs	fpn_convsappendfpn_bottleneck)rH   r4   r  rp   s      r-   r8   zBeitUperHead.__init__a  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 rL   )r:   rU   r>  r@  )rH   r   r0  s      r-   psp_forwardzBeitUperHead.psp_forward  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   rM   r   r   Fr   rN   )ziprA  rC  rF  r?  r   rP   r   r   r   rB  r:   rU   rD  r  )
rH   rG  lateralslateral_convr0  used_backbone_levelsr   
prev_shapefpn_outsoutputs
             r-   r`   zBeitUperHead.forward  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#   r8   rk   r:   r   rF  r`   r   r   s   @r-   r:  r:  Y  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.
    r4   in_indexr  r  r5   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 )Nr   r   )r  r  r  r"   r=  r*  )rg   r8   r<   r  auxiliary_channelsr(  auxiliary_num_convs	num_convsauxiliary_concat_inputconcat_inputrR  r   r   convsrC  r  r   conv_catr!  r   r  )rH   r4   rR  r  r  conv_paddingrW   rp   s          r-   r8   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,   rG  c                     || j                      }|}| j                  D ]
  } ||      } | j                  r(| j                  t	        j
                  ||gd            }| j                  |      }|S )Nr"   rN   )rR  rY  rX  rZ  r:   rU   r  )rH   rG  r   r   r-  s        r-   r`   zBeitFCNHead.forward  sn    (7 JJ 	0D /M	0 MM%))X}4MST*UVM6r,   )r   r   r"   )r'   r(   r)   r*   r#   r   rj   r8   rk   r:   r   r`   r   r   s   @r-   rQ  rQ    su     no!U !U,/!UBE!UUX[`adfiai[jUj!U	!UFT%,,-? ELL r,   rQ  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.r<   r  r  r5   Nc                     t         |           t        j                  ||||      | _        t        j
                  |      | _        t        j                         | _        t        j                  ||||      | _	        y )Nr  r  )
rg   r8   r   ConvTranspose2dconv_transpose1BatchNorm2dnormalizationGELUr  conv_transpose2)rH   r<   r  r  rp   s       r-   r8   zBeitFPNUpBlock.__init__  sb    !11+{Xclrs^^K8'')!11+{Xclrsr,   r   c                     | j                  |      }| j                  |      }| j                  |      }| j                  |      }|S ra   )rb  rd  r  rf  r   s     r-   r`   zBeitFPNUpBlock.forward  sF    ,,];**=96,,];r,   )r   r   )
r'   r(   r)   r*   r   r8   r:   r   r`   r   r   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.
    r4   c                     t         |           t        |j                        | _        t        j                  |j                  |j                  dd      | _        t        j                  dd      | _	        y )Nr   r`  )
rg   r8   r^  r<   fpn1r   ra  fpn2	MaxPool2dfpn4r   s     r-   r8   zBeitFPNNeck.__init__  sX    "6#5#56	&&v'9'96;M;M[\efg	LLQq9	r,   feature_maps.r5   c                     | j                  |d         | j                  |d         |d   | j                  |d         fS rf   )rk  rl  rn  )rH   ro  s     r-   r`   zBeitFPNNeck.forward  sC    IIl1o&IIl1o&OIIl1o&	
 	
r,   )r'   r(   r)   r*   r#   r8   rj   r:   r   r`   r   r   s   @r-   ri  ri    sD    
:z :
E%,,*;$< 
u||UXGXAY 
r,   ri  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 )BeitForSemanticSegmentationr4   r5   Nc                 `   t         |   |       |j                  | _        t        |d      | _        t        | j                  j                        dk7  rt        d      t        |      | _
        t        |      | _        |j                  rt        |      nd | _        | j!                          y )NFr   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.)rg   r8   r   r   r   r?  r4   out_indices
ValueErrorri  fpnr:  decode_headuse_auxiliary_headrQ  auxiliary_headr   r   s     r-   r8   z$BeitForSemanticSegmentation.__init__  s      ++f>	t{{&&'1,- 
 v& (/5;5N5Nk&1TX 	r,   rI   r  rV   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 onerV   c              3      K   | ]7  }|d z
     ddd df   j                  d d      j                  d       9 yw)r"   Nr   rM   )	transposer   ).0r   r[   rG  patch_heightpatch_widths     r-   	<genexpr>z6BeitForSemanticSegmentation.forward.<locals>.<genexpr>U  sN      
 "!a%(AB/99!Q?GG
TVXdfqr
s   =A )ignore_indexauxiliary_logitsauxiliary_loss_weightr  )r4   r   ru  r   r   rP   rA   rj   rt  rv  rw  ry  r  semantic_loss_ignore_indexr  r   r	  )rH   rI   r  rV   r   r  rW   rX   rY   ro  r  r  r  r[   rG  r~  r  s                @@@@r-   r`   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,   r  )r'   r(   r)   r#   r8   r   r	   r   r:   r   r   r   r   rj   r   r`   r   r   s   @r-   rr  rr    s    z d *   -1&*).	G
llT)G
 t#G
 #'	G

 +,G
 
(	(G
  ! G
r,   rr  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"   Fr   )rg   r8   r   r   r<   num_featuresr   r   add_fpnri  r   r   rv  r   )rH   r4   rW   rp   s      r-   r8   zBeitBackbone.__init__x  st     9>v?W?WZ[?[9\]AV//]f>	*0..;v&bkkm 	 ^s   BrI   r   r5   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"   r   rM   )ro  r   r	  )rP   r4   rA   r   r   rI  stage_namesout_featuresreshape_hidden_statesr|  r   rv  r   r	  )rH   rI   r   r[   rW   rX   rY   r~  r  r  r   ro  stager0  s                 r-   r`   zBeitBackbone.forward  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)   r8   r   r	   r   r   r   r   r   r`   r   r   s   @r-   r  r  r  sN      3
3
 +,3
 
	3
  ! 3
r,   r  )r  r   rr  r   r   r  )Hr*   dataclassesr   r:   r   r    r   r   backbone_utilsr   r	   masking_utilsr
   modeling_outputsr   r   r   r   r   modeling_utilsr   processing_utilsr   pytorch_utilsr   utilsr   r   r   utils.genericr   r   utils.output_capturingr   resnet.modeling_resnetr   swin.modeling_swinr   vit.modeling_vitr   r   r   r   r    r!   configuration_beitr#   r&   r/   r3   r7   rd   r   r   r   r   r   r   r   r   r  r  r&  r2  r:  rQ  r^  ri  rr  r  __all__r+   r,   r-   <module>r     sF    !   & H 6  . & @ B B I 5 4 - t t * 
 !;  	, 	'] 'TV3ryy V3r`L `	f 		< 	7 7t T, T T. Jj# Jj JjZ	v 	v S
!4 S
S
l /
!4 /
/
d
O 
4
bii 
Sryy S<N299 Nb:")) :zRYY $
")) 
* `
"5 `
 `
F 
A
="5 A

A
Hr,   