Ë
    S^(h!~  ã            	       óæ  — d Z ddlZddlZddlmZmZ ddlZddlZddl	Zddlm
Z
mZ ddlmZmZmZ ddlmZ ddlmZmZ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"  ejF                  e$«      Z%dZ&dZ'g d¢Z(dZ)dZ*d=deee+f   fd„Z, G d„ dejZ                  «      Z. G d„ dej^                  «      Z0 G d„ dejb                  «      Z2 G d„ dejf                  «      Z4 G d„ dejb                  «      Z5d>dej                  de6d e+dej                  fd!„Z7 G d"„ d#ejb                  «      Z8d?d$„Z9 G d%„ d&ejb                  «      Z: G d'„ d(ejb                  «      Z; G d)„ d*ejb                  «      Z< G d+„ d,ejb                  «      Z= G d-„ d.ejb                  «      Z> G d/„ d0e«      Z?d1Z@d2ZA ed3e@«       G d4„ d5e?«      «       ZB ed6e@«       G d7„ d8e?«      «       ZC ed9e@«       G d:„ d;e?e «      «       ZDg d<¢ZEy)@z9PyTorch BiT model. Also supports backbone for ViT hybrid.é    N)ÚOptionalÚTuple)ÚTensorÚnn)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELossé   )ÚACT2FN)ÚBackboneOutputÚBaseModelOutputWithNoAttentionÚ(BaseModelOutputWithPoolingAndNoAttentionÚ$ImageClassifierOutputWithNoAttention)ÚPreTrainedModel)Úadd_code_sample_docstringsÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingÚreplace_return_docstrings)ÚBackboneMixiné   )Ú	BitConfigr   zgoogle/bit-50)r   i   é   r   z	tiger catÚreturnc                 ó  — d}| €|dz
  ||dz
  z  z   dz  } | |fS t        | t        «      ra| j                  «       } | dk(  r0|dk(  r#||dz
  z  dz  dk(  r|dz
  ||dz
  z  z   dz  } | |fS d} d}| |fS | dk(  rd} | |fS |dz
  ||dz
  z  z   dz  } | |fS )al  
    Utility function to get the tuple padding value given the kernel_size and padding.

    Args:
        padding (Union[`str`, `int`], *optional*):
            Padding value, can be either `"same"`, `"valid"`. If a different value is provided the default padding from
            PyTorch is used.
        kernel_size (`int`, *optional*, defaults to 7):
            Kernel size of the convolution layers.
        stride (`int`, *optional*, defaults to 1):
            Stride value of the convolution layers.
        dilation (`int`, *optional*, defaults to 1):
            Dilation value of the convolution layers.
    Fr   é   Úsamer   TÚvalid)Ú
isinstanceÚstrÚlower)ÚpaddingÚkernel_sizeÚstrideÚdilationÚdynamics        úb/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/bit/modeling_bit.pyÚget_padding_valuer(   <   sõ   € ð €GØ€Ø˜Q‘J (¨k¸A©oÑ">Ñ>À1ÑDˆØ˜ÐÐä�'œ3Ôà—-‘-“/ˆØ�fÒà˜Š{ ¨K¸!©OÑ <ÀÑAÀQÒFà" Q™J¨(°kÀA±oÑ*FÑFÈ1ÑL�ð �GÐÐð �Ø�ð �GÐÐð ˜ÒàˆGð �GÐÐð  ™
 h°+À±/Ñ&BÑBÀqÑHˆGØ�GÐÐó    c                   ó6   ‡ — e Zd ZdZ	 	 	 	 	 	 dˆ fd„	Zd„ Zˆ xZS )ÚWeightStandardizedConv2dz÷Conv2d with Weight Standardization. Includes TensorFlow compatible SAME padding. Used for ViT Hybrid model.

    Paper: [Micro-Batch Training with Batch-Channel Normalization and Weight
    Standardization](https://arxiv.org/abs/1903.10520v2)
    c
           
      ó¬   •— t        ||||¬«      \  }}
t        ‰| �	  ||||||||¬«       |
rt        |||«      | _        |	| _        y d | _        |	| _        y )N)r$   r%   )r$   r"   r%   ÚgroupsÚbias)r(   ÚsuperÚ__init__ÚDynamicPad2dÚpadÚeps)ÚselfÚ
in_channelÚout_channelsr#   r$   r"   r%   r-   r.   r3   Ú
is_dynamicÚ	__class__s              €r'   r0   z!WeightStandardizedConv2d.__init__l   ss   ø€ ô 0°¸ÈVÐ^fÔgÑˆ�Ü‰ÑØØØØØØØØð 	ô 		
ñ Ü# K°¸ÓBˆDŒHð ˆ�ð ˆDŒHØˆ�r)   c           	      óÈ  — | j                   �| j                  |«      }t        j                  j                  | j                  j                  d| j                  d«      d d dd| j                  ¬«      j                  | j                  «      }t        j                  j                  ||| j                  | j                  | j                  | j                  | j                  «      }|S )Nr   éÿÿÿÿTç        )ÚtrainingÚmomentumr3   )r2   r   Ú
functionalÚ
batch_normÚweightÚreshaper6   r3   Ú
reshape_asÚconv2dr.   r$   r"   r%   r-   )r4   Úhidden_stater@   s      r'   Úforwardz WeightStandardizedConv2d.forward‰   s·   € Ø�8‰8ÐØŸ8™8 LÓ1ˆLÜ—‘×)Ñ)Ø�K‰K×Ñ  4×#4Ñ#4°bÓ9¸4ÀÐPTÐ_bÐhl×hpÑhpð *ó 
ç
‰*�T—[‘[Ó
!ð 	ô —}‘}×+Ñ+Ø˜& $§)¡)¨T¯[©[¸$¿,¹,ÈÏÉÐW[×WbÑWbó
ˆð Ðr)   )r   ÚSAMEr   r   Fg�íµ ÷Æ°>©Ú__name__Ú
__module__Ú__qualname__Ú__doc__r0   rE   Ú__classcell__©r8   s   @r'   r+   r+   e   s&   ø„ ñð ØØØØØõö:	r)   r+   c                   ó*   ‡ — e Zd ZdZdˆ fd„	Zd„ Zˆ xZS )ÚBitGroupNormActivationzQ
    A module that combines group normalization with an activation function.
    c                 ó°   •— t         t        | �  |j                  |||¬«       |rt        |j
                     | _        y t        j                  «       | _        y )N)r3   Úaffine)	r/   rO   r0   Ú
num_groupsr   Ú
hidden_actÚ
activationr   ÚIdentity)r4   ÚconfigÚnum_channelsr3   rQ   Úapply_activationr8   s         €r'   r0   zBitGroupNormActivation.__init__š   sF   ø€ ÜÔ$ dÑ4°V×5FÑ5FÈÐZ]ÐflÐ4ÔmÙÜ$ V×%6Ñ%6Ñ7ˆD�Oä Ÿk™k›mˆD�Or)   c                 ó¾   — t         j                  j                  || j                  | j                  | j
                  | j                  «      }| j                  |«      }|S ©N)r   r>   Ú
group_normrR   r@   r.   r3   rT   )r4   rD   s     r'   rE   zBitGroupNormActivation.forward¡   sH   € Ü—}‘}×/Ñ/°¸d¿o¹oÈtÏ{É{Ð\`×\eÑ\eÐgk×goÑgoÓpˆØ—‘ |Ó4ˆØÐr)   )gñhãˆµøä>TTrG   rM   s   @r'   rO   rO   •   s   ø„ ñõ,ör)   rO   c                   ó*   ‡ — e Zd ZdZdˆ fd„	Zd„ Zˆ xZS )r1   zŒ
    A module that wraps dynamic padding of any input, given the parameters of the convolutional layer and the input
    hidden states.
    c                 óæ   •— t         ‰| �  «        t        |t        «      r||f}t        |t        «      r||f}t        |t        «      r||f}|| _        || _        || _        || _        d„ }|| _        y )Nc                 óp   — t        t        j                  | |z  «      dz
  |z  |dz
  |z  z   dz   | z
  d«      S )Nr   r   )ÚmaxÚmathÚceil)Úxr#   r$   r%   s       r'   Úcompute_paddingz.DynamicPad2d.__init__.<locals>.compute_padding¾   sB   € ÜœŸ	™	 ! f¡*Ó-°Ñ1°VÑ;¸{ÈQ¹ÐRZÑ>ZÑZÐ]^Ñ^ÐabÑbÐdeÓfÐfr)   )	r/   r0   r   Úintr#   r$   r%   Úvaluerc   )r4   r#   r$   r%   re   rc   r8   s         €r'   r0   zDynamicPad2d.__init__­   sw   ø€ Ü‰ÑÔä�k¤3Ô'Ø&¨Ð4ˆKä�fœcÔ"Ø˜fÐ%ˆFä�h¤Ô$Ø  (Ð+ˆHà&ˆÔØˆŒØ ˆŒØˆŒ
ò	gð  /ˆÕr)   c           	      ó¶  — |j                  «       dd  \  }}| j                  || j                  d   | j                  d   | j                  d   «      }| j                  || j                  d   | j                  d   | j                  d   «      }|dkD  s|dkD  rBt
        j                  j                  ||dz  ||dz  z
  |dz  ||dz  z
  g| j                  ¬«      }|S )Néþÿÿÿr   r   r   )re   )	Úsizerc   r#   r$   r%   r   r>   r2   re   )r4   ÚinputÚinput_heightÚinput_widthÚpadding_heightÚpadding_widths         r'   rE   zDynamicPad2d.forwardÃ   sù   € à$)§J¡J£L°°Ð$5Ñ!ˆ�kð ×-Ñ-¨l¸D×<LÑ<LÈQÑ<OÐQU×Q\ÑQ\Ð]^ÑQ_Ðae×anÑanÐopÑaqÓrˆØ×,Ñ,¨[¸$×:JÑ:JÈ1Ñ:MÈtÏ{É{Ð[\É~Ð_c×_lÑ_lÐmnÑ_oÓpˆð ˜AÒ °Ò!2Ü—M‘M×%Ñ%Øà! QÑ&Ø! M°QÑ$6Ñ6Ø" aÑ'Ø" ^°qÑ%8Ñ8ð	ð —j‘jð &ó 	ˆEð ˆr)   )r   rG   rM   s   @r'   r1   r1   §   s   ø„ ñõ
/ö,r)   r1   c                   ó<   ‡ — e Zd ZdZ	 	 	 	 	 	 ddefˆ fd„Zd„ Zˆ xZS )ÚBitMaxPool2dz1Tensorflow like 'SAME' wrapper for 2D max poolingr#   c                 ó†  •— t        |t        j                  j                  «      r|n||f}t        |t        j                  j                  «      r|n||f}t        |t        j                  j                  «      r|n||f}t        ‰| �  |||||«       |rt        ||||«      | _        y t        j                  «       | _        y rZ   )
r   ÚcollectionsÚabcÚIterabler/   r0   r1   r2   r   rU   )	r4   r#   r$   r%   Ú	ceil_moder"   Úpadding_valueÚuse_dynamic_paddingr8   s	           €r'   r0   zBitMaxPool2d.__init__Ý   sž   ø€ ô &0°¼[¿_¹_×=UÑ=UÔ%V‘kÐ]hÐjuÐ\vˆÜ% f¬k¯o©o×.FÑ.FÔG‘ÈfÐV\ÐM]ˆÜ)¨(´K·O±O×4LÑ4LÔM‘8ÐT\Ð^fÐSgˆÜ‰Ñ˜ f¨g°xÀÔKÙÜ# K°¸À=ÓQˆD�Hä—{‘{“}ˆD�Hr)   c                 óÐ   — | j                  |«      }t        j                  j                  || j                  | j
                  | j                  | j                  | j                  «      S rZ   )	r2   r   r>   Ú
max_pool2dr#   r$   r"   r%   rt   ©r4   Úhidden_statess     r'   rE   zBitMaxPool2d.forwardð   sM   € ØŸ™ Ó/ˆÜ�}‰}×'Ñ'Ø˜4×+Ñ+¨T¯[©[¸$¿,¹,ÈÏÉÐW[×WeÑWeó
ð 	
r)   )Nr   F)r   r   r   T)rH   rI   rJ   rK   rd   r0   rE   rL   rM   s   @r'   ro   ro   Ú   s,   ø„ Ù;ð
 ØØØØØ ñ%àõ%ö&
r)   ro   c                   ó8   ‡ — e Zd ZdZdefˆ fd„Zdedefd„Zˆ xZS )ÚBitEmbeddingszL
    BiT Embeddings (stem) composed of a single aggressive convolution.
    rV   c                 ó.  •— t         ‰| �  «        t        |j                  |j                  ddd|j
                  ¬«      | _        t        dd|j                  ¬«      | _	        |j
                  �7|j
                  j                  «       dk(  rt        j                  «       | _        nt        j                  dd	¬
«      | _        |j                  dk(  st!        ||j                  ¬«      | _        nt        j                  «       | _        |j                  | _        y )Nr   r   ç:Œ0âŽyE>)r#   r$   r3   r"   r
   )r#   r$   rv   rF   )r   r   r   r   r;   )r"   re   Úpreactivation©rW   )r/   r0   r+   rW   Úembedding_sizeÚglobal_paddingÚconvolutionro   Úembedding_dynamic_paddingÚpoolerÚupperr   rU   r2   ÚConstantPad2dÚ
layer_typerO   Únorm©r4   rV   r8   s     €r'   r0   zBitEmbeddings.__init__ü   sÛ   ø€ Ü‰ÑÔä3Ø×ÑØ×!Ñ!ØØØØ×)Ñ)ô
ˆÔô #¨q¸ÐPV×PpÑPpÔqˆŒð × Ñ Ð,°×1FÑ1F×1LÑ1LÓ1NÐRXÒ1XÜ—{‘{“}ˆD�Hä×'Ñ'°ÀCÔHˆDŒHà× Ñ  OÒ3Ü.¨vÀF×DYÑDYÔZˆD�IäŸ™›ˆDŒIà"×/Ñ/ˆÕr)   Úpixel_valuesr   c                 óà   — |j                   d   }|| j                  k7  rt        d«      ‚| j                  |«      }| j	                  |«      }| j                  |«      }| j                  |«      }|S )Nr   zeMake sure that the channel dimension of the pixel values match with the one set in the configuration.)ÚshaperW   Ú
ValueErrorrƒ   r2   r‰   r…   )r4   r‹   rW   Ú	embeddings       r'   rE   zBitEmbeddings.forward  sr   € Ø#×)Ñ)¨!Ñ,ˆØ˜4×,Ñ,Ò,ÜØwóð ð ×$Ñ$ \Ó2ˆ	à—H‘H˜YÓ'ˆ	à—I‘I˜iÓ(ˆ	à—K‘K 	Ó*ˆ	àÐr)   )	rH   rI   rJ   rK   r   r0   r   rE   rL   rM   s   @r'   r|   r|   ÷   s'   ø„ ñð0˜yõ 0ð6 Fð ¨v÷ r)   r|   ri   Ú	drop_probr<   c                 ó  — |dk(  s|s| S d|z
  }| j                   d   fd| j                  dz
  z  z   }|t        j                  || j                  | j
                  ¬«      z   }|j                  «        | j                  |«      |z  }|S )aF  
    Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).

    Comment by Ross Wightman: This is the same as the DropConnect impl I created for EfficientNet, etc networks,
    however, the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
    See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for changing the
    layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use 'survival rate' as the
    argument.
    r;   r   r   )r   )ÚdtypeÚdevice)r�   ÚndimÚtorchÚrandr’   r“   Úfloor_Údiv)ri   r�   r<   Ú	keep_probr�   Úrandom_tensorÚoutputs          r'   Ú	drop_pathrœ   *  s�   € ð �CÒ™xØˆØ�I‘€IØ�[‰[˜‰^Ð ¨¯
©
°Q©Ñ 7Ñ7€EØ¤§
¡
¨5¸¿¹ÈEÏLÉLÔ YÑY€MØ×ÑÔØ�Y‰Y�yÓ! MÑ1€FØ€Mr)   c                   óx   ‡ — e Zd ZdZd	dee   ddfˆ fd„Zdej                  dej                  fd„Z	de
fd„Zˆ xZS )
ÚBitDropPathzXDrop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).Nr�   r   c                 ó0   •— t         ‰| �  «        || _        y rZ   )r/   r0   r�   )r4   r�   r8   s     €r'   r0   zBitDropPath.__init__B  s   ø€ Ü‰ÑÔØ"ˆ�r)   rz   c                 óD   — t        || j                  | j                  «      S rZ   )rœ   r�   r<   ry   s     r'   rE   zBitDropPath.forwardF  s   € Ü˜¨¯©¸¿¹ÓFÐFr)   c                 ó8   — dj                  | j                  «      S )Nzp={})Úformatr�   )r4   s    r'   Ú
extra_reprzBitDropPath.extra_reprI  s   € Ø�}‰}˜TŸ^™^Ó,Ð,r)   rZ   )rH   rI   rJ   rK   r   Úfloatr0   r•   r   rE   r    r£   rL   rM   s   @r'   rž   rž   ?  sG   ø„ Ùbñ# (¨5¡/ð #¸Tõ #ðG U§\¡\ð G°e·l±ló Gð-˜C÷ -r)   rž   c                 óf   — |}t        |t        | |dz  z   «      |z  |z  «      }|d| z  k  r||z  }|S )Nr   gÍÌÌÌÌÌì?)r_   rd   )re   ÚdivisorÚ	min_valueÚ	new_values       r'   Úmake_divr©   M  sG   € Ø€IÜ�Iœs 5¨7°Q©;Ñ#6Ó7¸7ÑBÀWÑLÓM€IØ�3˜‘;ÒØ�WÑˆ	ØÐr)   c                   ó:   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 dˆ fd„	Zd„ Zˆ xZS )ÚBitPreActivationBottleneckLayera  Pre-activation (v2) bottleneck block.
    Follows the implementation of "Identity Mappings in Deep Residual Networks":
    https://github.com/KaimingHe/resnet-1k-layers/blob/master/resnet-pre-act.lua

    Except it puts the stride on 3x3 conv when available.
    c           	      ó  •— t         ‰| �  «        |xs |}|xs |}t        ||z  «      }|
rt        ||||d¬«      | _        nd | _        t        ||«      | _        t        ||dd|j                  ¬«      | _	        t        ||¬«      | _
        t        ||d||d|j                  ¬«      | _        t        ||«      | _        t        ||dd|j                  ¬«      | _        |	d	kD  rt        |	«      | _        y t        j                   «       | _        y )
NT©r$   Úpreactr   r~   ©r3   r"   r€   r
   )r$   r-   r3   r"   r   )r/   r0   r©   ÚBitDownsampleConvÚ
downsamplerO   Únorm1r+   r‚   Úconv1Únorm2Úconv2Únorm3Úconv3rž   r   rU   rœ   )r4   rV   Úin_channelsr6   Úbottle_ratior$   r%   Úfirst_dilationr-   Údrop_path_rateÚis_first_layerÚmid_channelsr8   s               €r'   r0   z(BitPreActivationBottleneckLayer.__init__]  s  ø€ ô 	‰ÑÔà'Ò3¨8ˆà#Ò2 {ˆÜ ¨|Ñ ;Ó<ˆáÜ/ØØØØØôˆD�Oð #ˆDŒOä+¨F°KÓ@ˆŒ
Ü-¨k¸<ÈÐPTÐ^d×^sÑ^sÔtˆŒ
ä+¨FÀÔNˆŒ
Ü-Ø˜,¨°&ÀÈTÐ[a×[pÑ[pô
ˆŒ
ô ,¨F°LÓAˆŒ
Ü-¨l¸LÈ!ÐQUÐ_e×_tÑ_tÔuˆŒ
à8FÈÒ8Jœ ^Ó4ˆ�ÔPR×P[ÑP[ÓP]ˆ�r)   c                 ó0  — | j                  |«      }|}| j                  �| j                  |«      }| j                  |«      }| j                  | j	                  |«      «      }| j                  | j                  |«      «      }| j                  |«      }||z   S rZ   )r²   r±   r³   rµ   r´   r·   r¶   rœ   )r4   rz   Úhidden_states_preactÚshortcuts       r'   rE   z'BitPreActivationBottleneckLayer.forward‰  s‰   € Ø#Ÿz™z¨-Ó8Ðð !ˆØ�?‰?Ð&Ø—‘Ð';Ó<ˆHð Ÿ
™
Ð#7Ó8ˆØŸ
™
 4§:¡:¨mÓ#<Ó=ˆØŸ
™
 4§:¡:¨mÓ#<Ó=ˆØŸ™ }Ó5ˆØ˜xÑ'Ð'r)   ©Nç      Ð?r   r   Nr   r;   FrG   rM   s   @r'   r«   r«   U  s.   ø„ ñð ØØØØØØØõ*^öX(r)   r«   c                   ó:   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 dˆ fd„	Zd„ Zˆ xZS )ÚBitBottleneckLayerz\Non Pre-activation bottleneck block, equivalent to V1.5/V1b bottleneck. Used for ViT Hybrid.c           
      óD  •— t         ‰| �  «        |xs |}|xs |}t        ||z  «      }|
rt        ||||d¬«      | _        nd | _        t        ||dd|j                  ¬«      | _        t        ||¬«      | _	        t        ||d|||d|j                  ¬«      | _
        t        ||¬«      | _        t        ||dd|j                  ¬«      | _        t        ||d¬	«      | _        |	d
kD  rt        |	«      nt        j                   «       | _        t$        |j&                     | _        y )NFr­   r   r~   r¯   r€   r
   )r$   r%   r-   r3   r"   ©rW   rX   r   )r/   r0   r©   r°   r±   r+   r‚   r³   rO   r²   rµ   r´   r·   r¶   rž   r   rU   rœ   r   rS   rT   )r4   rV   r¸   r6   r¹   r$   r%   rº   r-   r»   r¼   Úmid_chsr8   s               €r'   r0   zBitBottleneckLayer.__init__œ  s  ø€ ô 	‰ÑÔØ'Ò3¨8ˆà#Ò2 {ˆÜ˜<¨,Ñ6Ó7ˆáÜ/ØØØØØôˆD�Oð #ˆDŒOä-¨k¸7ÀAÈ4ÐY_×YnÑYnÔoˆŒ
Ü+¨FÀÔIˆŒ
Ü-ØØØØØ#ØØØ×)Ñ)ô	
ˆŒ
ô ,¨FÀÔIˆŒ
Ü-¨g°|ÀQÈDÐZ`×ZoÑZoÔpˆŒ
Ü+¨FÀÐ`eÔfˆŒ
Ø8FÈÒ8Jœ ^Ô4ÔPR×P[ÑP[ÓP]ˆŒä  ×!2Ñ!2Ñ3ˆ�r)   c                 óZ  — |}| j                   �| j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j	                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j                  ||z   «      }|S rZ   )	r±   r³   r²   rµ   r´   r·   r¶   rœ   rT   )r4   rz   rÀ   s      r'   rE   zBitBottleneckLayer.forwardÍ  sœ   € à ˆØ�?‰?Ð&Ø—‘ }Ó5ˆHð Ÿ
™
 =Ó1ˆØŸ
™
 =Ó1ˆàŸ
™
 =Ó1ˆØŸ
™
 =Ó1ˆàŸ
™
 =Ó1ˆØŸ
™
 =Ó1ˆàŸ™ }Ó5ˆØŸ™¨¸Ñ(@ÓAˆØÐr)   rÁ   rG   rM   s   @r'   rÄ   rÄ   ™  s+   ø„ Ùfð ØØØØØØØõ/4öbr)   rÄ   c                   ó*   ‡ — e Zd Z	 	 dˆ fd„	Zd„ Zˆ xZS )r°   c                 óÀ   •— t         ‰| �  «        t        ||d|d|j                  ¬«      | _        |rt        j                  «       | _        y t        ||d¬«      | _        y )Nr   r~   )r$   r3   r"   FrÆ   )	r/   r0   r+   r‚   Úconvr   rU   rO   r‰   )r4   rV   r¸   r6   r$   r®   r8   s         €r'   r0   zBitDownsampleConv.__init__ã  s\   ø€ ô 	‰ÑÔÜ,Ø˜ q°¸TÈ6×K`ÑK`ô
ˆŒ	ñ
 ô �K‰K‹Mð 	�	ô (¨¸\Ð\aÔbð 	�	r)   c                 óB   — | j                  | j                  |«      «      S rZ   )r‰   rË   )r4   rb   s     r'   rE   zBitDownsampleConv.forwardõ  s   € Ø�y‰y˜Ÿ™ 1›Ó&Ð&r)   )r   T)rH   rI   rJ   r0   rE   rL   rM   s   @r'   r°   r°   â  s   ø„ ð Øõ
ö$'r)   r°   c                   ó>   ‡ — e Zd ZdZ	 	 dˆ fd„	Zd„ Zdedefd„Zˆ xZS )ÚBitStagez7
    A ResNet v2 stage composed by stacked layers.
    c	                 ó^  •— t         ‰| �  «        |dv rdnd}	|j                  dk(  rt        }
nt        }
|}t        j                  «       | _        t        |«      D ]Q  }| j                  |||«      \  }}}| j                  j                  t        |«       |
|||||||	||¬«	      «       |}|}	ŒS y )N)r   r   r   r   Ú
bottleneck)r$   r%   r¹   rº   r»   r¼   )r/   r0   rˆ   rÄ   r«   r   Ú
SequentialÚlayersÚrangeÚ_get_updated_hyperparametersÚ
add_moduler    )r4   rV   r¸   r6   r$   r%   Údepthr¹   Úlayer_dropoutrº   Ú	layer_clsÚprev_chsÚ	layer_idxr»   r¼   r8   s                  €r'   r0   zBitStage.__init__þ  sÅ   ø€ ô 	‰ÑÔà&¨&Ñ0™°aˆð ×Ñ Ò,Ü*‰Iä7ˆIàˆÜ—m‘m“oˆŒÜ˜u›ò 	&ˆIà59×5VÑ5VØ˜6 =ó6Ñ2ˆF�N Nð �K‰K×"Ñ"Ü�I“ÙØØØ Ø!Ø%Ø!-Ø#1Ø#1Ø#1ô
ôð $ˆHØ%‰Nñ+	&r)   c                 ó8   — |r||   }nd}|dk7  rd}|dk(  }|||fS )zt
        Get the new hyper-parameters with respect to the previous ones and the index of the current layer.
        r;   r   r   © )r4   rÚ   r$   r×   r»   r¼   s         r'   rÔ   z%BitStage._get_updated_hyperparameters,  s8   € ñ Ø*¨9Ñ5‰Nà ˆNà˜Š>ØˆFà" a™ˆà�~ ~Ð5Ð5r)   ri   r   c                 óT   — |}t        | j                  «      D ]  \  }} ||«      }Œ |S rZ   )Ú	enumeraterÒ   )r4   ri   rD   Ú_Úlayers        r'   rE   zBitStage.forward<  s3   € ØˆÜ! $§+¡+Ó.ò 	/‰HˆAˆuÙ  Ó.‰Lð	/àÐr)   )rÂ   N)	rH   rI   rJ   rK   r0   rÔ   r   rE   rL   rM   s   @r'   rÎ   rÎ   ù  s.   ø„ ñð Øõ,&ò\6ð ˜Vð ¨÷ r)   rÎ   c            	       óF   ‡ — e Zd Zdefˆ fd„Zd„ Z	 d	dedededefd„Z	ˆ xZ
S )
Ú
BitEncoderrV   c           
      ó�  •— t         ‰| �  «        t        j                  g «      | _        |j
                  }d}d}t        j                  t        j                  d|j                  t        |j                  «      «      «      j                  |j                  «      D �cg c]  }|j                  «       ‘Œ }}t        t!        |j                  |j"                  |«      «      D ]`  \  }\  }}	}
| j%                  |||	||«      \  }}}t'        |||||||
¬«      }|}||z  }| j                  j)                  t+        |«      |«       Œb y c c}w )Né   r   r   )r$   r%   rÖ   r×   )r/   r0   r   Ú
ModuleListÚstagesr�   r•   r   ÚnpÚlinspacer»   ÚsumÚdepthsÚsplitÚtolistrÞ   ÚzipÚhidden_sizesrÔ   rÎ   rÕ   r    )r4   rV   rÙ   Úcurrent_strider%   rb   Úlayer_dropoutsÚ	stage_idxÚcurrent_depthÚcurrent_hidden_sizer×   r6   r$   Ústager8   s                 €r'   r0   zBitEncoder.__init__D  sA  ø€ Ü‰ÑÔÜ—m‘m BÓ'ˆŒà×(Ñ(ˆð ˆØˆô —\‘\¤"§+¡+¨a°×1FÑ1FÌÈFÏMÉMÓHZÓ"[Ó\×bÑbÐci×cpÑcpÓqö
àð �H‰H�Jð
ˆð 
ô
 OXÜ�—‘˜v×2Ñ2°NÓCóO
ò 	:ÑJˆIÑJ˜Ð':¸Mð .2×-NÑ-NØ˜>Ð+>ÀÈ&ó.Ñ*ˆL˜& (ô ØØØØØ!Ø#Ø+ôˆEð $ˆHØ˜fÑ$ˆNà�K‰K×"Ñ"¤3 y£>°5Õ9ñ+	:ùò
s   ÂEc                 óz   — t        ||j                  z  «      }|dk(  rdnd}||j                  k\  r||z  }d}|||fS )Nr   r   r   )r©   Úwidth_factorÚoutput_stride)r4   rñ   rï   ró   r%   rV   r6   r$   s           r'   rÔ   z'BitEncoder._get_updated_hyperparametersj  sO   € ÜÐ 3°f×6IÑ6IÑ IÓJˆØ 1’n‘¨!ˆØ˜V×1Ñ1Ò1Ø˜ÑˆHØˆFØ˜V XÐ-Ð-r)   rD   Úoutput_hidden_statesÚreturn_dictr   c                 ó¦   — |rdnd }| j                   D ]  }|r||fz   } ||«      }Œ |r||fz   }|st        d„ ||fD «       «      S t        ||¬«      S )NrÜ   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wrZ   rÜ   )Ú.0Úvs     r'   ú	<genexpr>z%BitEncoder.forward.<locals>.<genexpr>�  s   è ø€ ÒS˜qÀQÁ]œÑSùs   ‚Š)Úlast_hidden_staterz   )ræ   Útupler   )r4   rD   rø   rù   rz   Ústage_modules         r'   rE   zBitEncoder.forwardr  sv   € ñ 3™¸ˆà ŸK™Kò 	6ˆLÙ#Ø -°°Ñ ?�á'¨Ó5‰Lð		6ñ  Ø)¨\¨OÑ;ˆMáÜÑS \°=Ð$AÔSÓSÐSä-Ø*Ø'ô
ð 	
r)   )FT)rH   rI   rJ   r   r0   rÔ   r   Úboolr   rE   rL   rM   s   @r'   râ   râ   C  sA   ø„ ð$:˜yõ $:òL.ð ]añ
Ø"ð
Ø:>ð
ØUYð
à	'÷
r)   râ   c                   ó(   — e Zd ZdZeZdZdZdgZd„ Z	y)ÚBitPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    Úbitr‹   r|   c                 óJ  — t        |t        j                  «      r-t        j                  j	                  |j
                  dd¬«       y t        |t        j                  «      rÃt        j                  j                  |j
                  t        j                  d«      ¬«       |j                  �xt        j                  j                  |j
                  «      \  }}|dkD  rdt        j                  |«      z  nd}t        j                  j                  |j                  | |«       y y t        |t        j                  t        j                  f«      rUt        j                  j                  |j
                  d«       t        j                  j                  |j                  d«       y y )NÚfan_outÚrelu)ÚmodeÚnonlinearityé   )Úar   r   )r   r   ÚConv2dÚinitÚkaiming_normal_r@   ÚLinearÚkaiming_uniform_r`   Úsqrtr.   Ú_calculate_fan_in_and_fan_outÚuniform_ÚBatchNorm2dÚ	GroupNormÚ	constant_)r4   ÚmoduleÚfan_inrß   Úbounds        r'   Ú_init_weightsz BitPreTrainedModel._init_weights”  s  € Ü�fœbŸi™iÔ(Ü�G‰G×#Ñ# F§M¡M¸	ÐPVÐ#ÕWä˜¤§	¡	Ô*Ü�G‰G×$Ñ$ V§]¡]´d·i±iÀ³lÐ$ÔCØ�{‰{Ð&ÜŸG™G×AÑAÀ&Ç-Á-ÓP‘	�˜Ø17¸!²˜œDŸI™I fÓ-Ò-À�Ü—‘× Ñ  §¡¨u¨f°eÕ<ð 'ô ˜¤§¡´·±Ð >Ô?Ü�G‰G×Ñ˜fŸm™m¨QÔ/Ü�G‰G×Ñ˜fŸk™k¨1Õ-ð @r)   N)
rH   rI   rJ   rK   r   Úconfig_classÚbase_model_prefixÚmain_input_nameÚ_no_split_modulesr  rÜ   r)   r'   r  r  ‰  s'   „ ñð
 €LØÐØ$€OØ(Ð)Ðó.r)   r  aE  
    This model is a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass. Use it
    as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage and
    behavior.

    Parameters:
        config ([`BitConfig`]): Model configuration class with all the parameters of the model.
            Initializing with a config file does not load the weights associated with the model, only the
            configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.
aA  
    Args:
        pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
            Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See [`BitImageProcessor.__call__`]
            for details.

        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
zLThe bare BiT model outputting raw features without any specific head on top.c                   ó|   ‡ — e Zd Zˆ fd„Z ee«       eeee	de
¬«      	 d	dedee   dee   defd„«       «       Zˆ xZS )
ÚBitModelc                 óJ  •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        |j                  dk(  rt        ||j                  d   ¬«      nt        j                  «       | _        t        j                  d«      | _        | j                  «        y )Nr   r:   r€   )r   r   )r/   r0   rV   r|   Úembedderrâ   Úencoderrˆ   rO   rî   r   rU   r‰   ÚAdaptiveAvgPool2dr…   Ú	post_initrŠ   s     €r'   r0   zBitModel.__init__Á  s„   ø€ Ü‰Ñ˜Ô ØˆŒä% fÓ-ˆŒä! &Ó)ˆŒð × Ñ  OÒ3ô # 6¸×8KÑ8KÈBÑ8OÕPä—‘“ð 	Œ	ô ×*Ñ*¨6Ó2ˆŒà�‰Õr)   Úvision)Ú
checkpointÚoutput_typer  ÚmodalityÚexpected_outputr‹   rø   rù   r   c                 óJ  — |�|n| j                   j                  }|�|n| j                   j                  }| j                  |«      }| j	                  |||¬«      }|d   }| j                  |«      }| j                  |«      }|s
||f|dd  z   S t        |||j                  ¬«      S )N©rø   rù   r   r   )rÿ   Úpooler_outputrz   )	rV   rø   Úuse_return_dictr#  r$  r‰   r…   r   rz   )r4   r‹   rø   rù   Úembedding_outputÚencoder_outputsrÿ   Úpooled_outputs           r'   rE   zBitModel.forwardÒ  sÃ   € ð %9Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàŸ=™=¨Ó6ÐàŸ,™,ØÐ3GÐU`ð 'ó 
ˆð ,¨AÑ.Ðà ŸI™IÐ&7Ó8ÐàŸ™Ð$5Ó6ˆáØ% }Ð5¸ÈÈÐ8KÑKÐKä7Ø/Ø'Ø)×7Ñ7ô
ð 	
r)   ©NN)rH   rI   rJ   r0   r   ÚBIT_INPUTS_DOCSTRINGr   Ú_CHECKPOINT_FOR_DOCr   Ú_CONFIG_FOR_DOCÚ_EXPECTED_OUTPUT_SHAPEr   r   r  rE   rL   rM   s   @r'   r!  r!  ¼  sp   ø„ ô
ñ" +Ð+?Ó@ÙØ&Ø<Ø$ØØ.ôð ptñ
Ø"ð
Ø:BÀ4¹.ð
Ø^fÐgkÑ^lð
à	1ò
óó Aô
r)   r!  zƒ
    BiT Model with an image classification head on top (a linear layer on top of the pooled features), e.g. for
    ImageNet.
    c                   ó¸   ‡ — e Zd Zˆ fd„Z ee«       eeee	e
¬«      	 	 	 	 d	deej                     deej                     dee   dee   def
d„«       «       Zˆ xZS )
ÚBitForImageClassificationc                 ó|  •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        j                  t        j                  «       |j                  dkD  r-t        j                  |j                  d   |j                  «      nt        j                  «       «      | _        | j                  «        y )Nr   r:   )r/   r0   Ú
num_labelsr!  r  r   rÑ   ÚFlattenr  rî   rU   Ú
classifierr&  rŠ   s     €r'   r0   z"BitForImageClassification.__init__   s‡   ø€ Ü‰Ñ˜Ô Ø ×+Ñ+ˆŒÜ˜FÓ#ˆŒäŸ-™-Ü�J‰J‹LØEK×EVÑEVÐYZÒEZŒB�I‰I�f×)Ñ)¨"Ñ-¨v×/@Ñ/@ÔAÔ`b×`kÑ`kÓ`mó
ˆŒð
 	�‰Õr)   )r(  r)  r  r+  r‹   Úlabelsrø   rù   r   c                 ó  — |�|n| j                   j                  }| j                  |||¬«      }|r|j                  n|d   }| j	                  |«      }d}|��‡| j                   j
                  €�| j                  dk(  rd| j                   _        nl| j                  dkD  rL|j                  t        j                  k(  s|j                  t        j                  k(  rd| j                   _        nd| j                   _        | j                   j
                  dk(  rIt        «       }	| j                  dk(  r& |	|j                  «       |j                  «       «      }nŒ |	||«      }n‚| j                   j
                  dk(  r=t        «       }	 |	|j                  d| j                  «      |j                  d«      «      }n,| j                   j
                  dk(  rt        «       }	 |	||«      }|s|f|dd z   }
|�|f|
z   S |
S t!        |||j"                  ¬	«      S )
a0  
        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 classification loss is computed (Cross-Entropy).
        Nr-  r   Ú
regressionÚsingle_label_classificationÚmulti_label_classificationr:   r   )ÚlossÚlogitsrz   )rV   r/  r  r.  r=  Úproblem_typer;  r’   r•   Úlongrd   r	   Úsqueezer   Úviewr   r   rz   )r4   r‹   r>  rø   rù   Úoutputsr2  rD  rC  Úloss_fctr›   s              r'   rE   z!BitForImageClassification.forward  s»  € ð& &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—(‘(˜<Ð>RÐ`k�(Ólˆá1<˜×-Ò-À'È!Á*ˆà—‘ Ó/ˆàˆàÑØ�{‰{×'Ñ'Ð/Ø—?‘? aÒ'Ø/;�D—K‘KÕ,Ø—_‘_ qÒ(¨f¯l©l¼e¿j¹jÒ.HÈFÏLÉLÔ\a×\eÑ\eÒLeØ/L�D—K‘KÕ,à/K�D—K‘KÔ,Ø�{‰{×'Ñ'¨<Ò7Ü"›9�Ø—?‘? aÒ'Ù# F§N¡NÓ$4°f·n±nÓ6FÓG‘Dá# F¨FÓ3‘DØ—‘×)Ñ)Ð-JÒJÜ+Ó-�Ù §¡¨B°·±Ó @À&Ç+Á+ÈbÃ/ÓR‘Ø—‘×)Ñ)Ð-IÒIÜ,Ó.�Ù ¨Ó/�áØ�Y ¨¨ Ñ,ˆFØ'+Ð'7�D�7˜VÑ#ÐC¸VÐCä3¸ÀfÐ\c×\qÑ\qÔrÐrr)   )NNNN)rH   rI   rJ   r0   r   r4  r   Ú_IMAGE_CLASS_CHECKPOINTr   r6  Ú_IMAGE_CLASS_EXPECTED_OUTPUTr   r•   ÚFloatTensorÚ
LongTensorr  rE   rL   rM   s   @r'   r9  r9  ø  sŸ   ø„ ô
ñ +Ð+?Ó@ÙØ*Ø8Ø$Ø4ô	ð 59Ø-1Ø/3Ø&*ñ/sà˜u×0Ñ0Ñ1ð/sð ˜×)Ñ)Ñ*ð/sð ' t™nð	/sð
 ˜d‘^ð/sð 
.ò/sóó Aô/sr)   r9  zL
    BiT backbone, to be used with frameworks like DETR and MaskFormer.
    c                   óv   ‡ — e Zd Zˆ fd„Z ee«       eee¬«      	 dde	de
e   de
e   defd„«       «       Zˆ xZS )	ÚBitBackbonec                 óÀ   •— t         ‰| �  |«       t         ‰| �	  |«       t        |«      | _        |j
                  g|j                  z   | _        | j                  «        y rZ   )	r/   r0   Ú_init_backboner!  r  r�   rî   Únum_featuresr&  rŠ   s     €r'   r0   zBitBackbone.__init__L  sQ   ø€ Ü‰Ñ˜Ô Ü‰Ñ˜vÔ&ä˜FÓ#ˆŒØ#×2Ñ2Ð3°f×6IÑ6IÑIˆÔð 	�‰Õr)   )r)  r  r‹   rø   rù   r   c                 óŽ  — |�|n| j                   j                  }|�|n| j                   j                  }| j                  |dd¬«      }|j                  }d}t        | j                  «      D ]  \  }}|| j                  v sŒ|||   fz  }Œ |s|f}	|r|	|j                  fz  }	|	S t        ||r|j                  d¬«      S dd¬«      S )a`  
        Returns:

        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("google/bit-50")
        >>> model = AutoBackbone.from_pretrained("google/bit-50")

        >>> inputs = processor(image, return_tensors="pt")
        >>> outputs = model(**inputs)
        ```NTr-  rÜ   )Úfeature_mapsrz   Ú
attentions)	rV   r/  rø   r  rz   rÞ   Ústage_namesÚout_featuresr   )
r4   r‹   rø   rù   rI  rz   rU  Úidxrô   r›   s
             r'   rE   zBitBackbone.forwardV  sî   € ð2 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð —(‘(˜<¸dÐPT�(ÓUˆà×-Ñ-ˆàˆÜ# D×$4Ñ$4Ó5ò 	6‰JˆC�Ø˜×)Ñ)Ò)Ø ¨sÑ!3Ð 5Ñ5‘ð	6ñ Ø"�_ˆFÙ#Ø˜7×0Ñ0Ð2Ñ2�ØˆMäØ%Ù3G˜'×/Ñ/Øô
ð 	
àMQØô
ð 	
r)   r3  )rH   rI   rJ   r0   r   r4  r   r   r6  r   r   r  rE   rL   rM   s   @r'   rP  rP  E  s`   ø„ ôñ +Ð+?Ó@Ù¨>ÈÔXàosñ/
Ø"ð/
Ø:BÀ4¹.ð/
Ø^fÐgkÑ^lð/
à	ò/
ó Yó Aô/
r)   rP  )r9  r!  r  rP  )Nr   r   r   )r;   F)é   )FrK   rq   r`   Útypingr   r   Únumpyrç   r•   Útorch.utils.checkpointr   r   Útorch.nnr   r   r	   Úactivationsr   Úmodeling_outputsr   r   r   r   Úmodeling_utilsr   Úutilsr   r   r   r   r   Úutils.backbone_utilsr   Úconfiguration_bitr   Ú
get_loggerrH   Úloggerr6  r5  r7  rK  rL  r  r(   r  r+   r  rO   ÚModuler1   Ú	MaxPool2dro   r|   r¤   rœ   rž   r©   r«   rÄ   r°   rÎ   râ   r  ÚBIT_START_DOCSTRINGr4  r!  r9  rP  Ú__all__rÜ   r)   r'   ú<module>rk     s6  ðñ @ã Û ß "ã Û Û ß ß AÑ Aå !÷ó õ .÷õ õ 2Ý (ð 
ˆ×	Ñ	˜HÓ	%€ð €ð &Ð Ú(Ð ð *Ð Ø*Ð ñ&ÈEÐRWÐY]ÐR]ÑL^ó &ôR-˜rŸy™yô -ô`˜RŸ\™\ô ô$0�2—9‘9ô 0ôf
�2—<‘<ô 
ô:/�B—I‘Iô /ñf�U—\‘\ð ¨eð ÀTð ÐV[×VbÑVbó ô*-�"—)‘)ô -óôA( b§i¡iô A(ôHF˜Ÿ™ô FôR'˜Ÿ	™	ô 'ô.Gˆr�y‰yô GôTC
�—‘ô C
ôL.˜ô .ð4	Ð ðÐ ñ ØRØóô5
Ð!ó 5
ó	ð5
ñp ðð óôCsÐ 2ó CsóðCsñL ðð ó	ô<
Ð$ mó <
óð<
ò~ Y�r)   