Ë
    þÍ:jn}  ã                   ó®  — d dl Z d dlmZ d dlmZ d dlmZ d dlmZm	Z	m
Z
 d dlZd dlZd dlmc m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 d d
lmZmZ d dlm Z  d dl!m"Z"m#Z# d dl$m%Z% g d¢Z&de'e(e(f   de(de(de(de'e(e(f   f
d„Z)de'e(e(f   de(de*e'e(e(f      fd„Z+de(de(dej$                  fd„Z, G d„ dejZ                  «      Z. G d„ dejZ                  «      Z/ G d„ dejZ                  «      Z0 G d „ d!ejZ                  «      Z1 G d"„ d#ejZ                  «      Z2 G d$„ d%ejZ                  «      Z3 G d&„ d'ejZ                  «      Z4 G d(„ d)ejZ                  «      Z5 G d*„ d+ejZ                  «      Z6	 	 d=d,e(d-e*e(   d.e*e(   d/e7d0e(d1e(d2e
e   d3e8d4ede6fd5„Z9 G d6„ d7e«      Z: e«        ed8e:jv                  f¬9«      dd:d;œd2e
e:   d3e8d4ede6fd<„«       «       Z<y)>é    N)ÚOrderedDict)ÚSequence)Úpartial)ÚAnyÚCallableÚOptional)ÚnnÚTensor)Úregister_modelÚWeightsÚWeightsEnum)Ú_IMAGENET_CATEGORIES)Ú_ovewrite_named_paramÚhandle_legacy_interface)ÚConv2dNormActivationÚSqueezeExcitation)ÚStochasticDepth)ÚImageClassificationÚInterpolationMode)Ú_log_api_usage_once)ÚMaxVitÚMaxVit_T_WeightsÚmaxvit_tÚ
input_sizeÚkernel_sizeÚstrideÚpaddingÚreturnc                 óR   — | d   |z
  d|z  z   |z  dz   | d   |z
  d|z  z   |z  dz   fS )Nr   é   é   © )r   r   r   r   s       ún/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchvision/models/maxvit.pyÚ_get_conv_output_shaper$      sJ   € à	�A‰˜Ñ	$ q¨7¡{Ñ	2°vÑ=ÀÑAØ	�A‰˜Ñ	$ q¨7¡{Ñ	2°vÑ=ÀÑAðð ó    Ún_blocksc                 ó„   — g }t        | ddd«      }t        |«      D ]!  }t        |ddd«      }|j                  |«       Œ# |S )zQUtil function to check that the input size is correct for a MaxVit configuration.é   r    r!   )r$   ÚrangeÚappend)r   r&   ÚshapesÚblock_input_shapeÚ_s        r#   Ú_make_block_input_shapesr.   !   sQ   € à€FÜ.¨z¸1¸aÀÓCÐÜ�8‹_ò )ˆÜ2Ð3DÀaÈÈAÓNÐØ�‰Ð'Õ(ð)ð €Mr%   ÚheightÚwidthc                 óø  — t        j                  t        j                  t        j                  | «      t        j                  |«      gd¬«      «      }t        j                  |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z
  z  cc<   |d d …d d …dfxx   |dz
  z  cc<   |d d …d d …dfxx   d|z  dz
  z  cc<   |j                  d«      S )NÚij)Úindexingr!   r    r   éÿÿÿÿ)ÚtorchÚstackÚmeshgridÚarangeÚflattenÚpermuteÚ
contiguousÚsum)r/   r0   ÚcoordsÚcoords_flatÚrelative_coordss        r#   Ú_get_relative_position_indexr@   +   sã   € Ü�[‰[œŸ™¬¯©°fÓ)=¼u¿|¹|ÈEÓ?RÐ(SÐ^bÔcÓd€FÜ—-‘- ¨Ó*€KØ!¢!¢Q¨ *Ñ-°ºA¸tÂQ¸JÑ0GÑG€OØ%×-Ñ-¨a°°AÓ6×AÑAÓC€OØ’A’q˜!�GÓ ¨¡
Ñ*ÓØ’A’q˜!�GÓ ¨¡	Ñ)ÓØ’A’q˜!�GÓ  E¡	¨A¡Ñ-ÓØ×Ñ˜rÓ"Ð"r%   c                   ó˜   ‡ — e Zd ZdZ	 ddededededededej                  f   d	edej                  f   d
eddfˆ fd„Z	de
de
fd„Zˆ xZS )ÚMBConva=  MBConv: Mobile Inverted Residual Bottleneck.

    Args:
        in_channels (int): Number of input channels.
        out_channels (int): Number of output channels.
        expansion_ratio (float): Expansion ratio in the bottleneck.
        squeeze_ratio (float): Squeeze ratio in the SE Layer.
        stride (int): Stride of the depthwise convolution.
        activation_layer (Callable[..., nn.Module]): Activation function.
        norm_layer (Callable[..., nn.Module]): Normalization function.
        p_stochastic_dropout (float): Probability of stochastic depth.
    Úin_channelsÚout_channelsÚexpansion_ratioÚsqueeze_ratior   Úactivation_layer.Ú
norm_layerÚp_stochastic_dropoutr   Nc	                 óÞ  •— t         ‰| �  «        |  |dk7  xs ||k7  }	|	rTt        j                  ||ddd¬«      g}
|dk(  rt        j                  d|d¬«      g|
z   }
t        j
                  |
Ž | _        nt        j                  «       | _        t        ||z  «      }t        ||z  «      }|rt        |d¬«      | _
        nt        j                  «       | _
        t        «       } ||«      |d	<   t        ||ddd
||d ¬«      |d<   t        ||d|d|||d ¬«	      |d<   t        ||t        j                  ¬«      |d<   t        j                  ||dd¬«      |d<   t        j
                  |«      | _        y )Nr!   T)r   r   Úbiasr    r(   ©r   r   r   Úrow©ÚmodeÚpre_normr   )r   r   r   rG   rH   ÚinplaceÚconv_a)r   r   r   rG   rH   ÚgroupsrQ   Úconv_b)Ú
activationÚsqueeze_excitation)rC   rD   r   rK   Úconv_c)ÚsuperÚ__init__r	   ÚConv2dÚ	AvgPool2dÚ
SequentialÚprojÚIdentityÚintr   Ústochastic_depthr   r   r   ÚSiLUÚlayers)ÚselfrC   rD   rE   rF   r   rG   rH   rI   Úshould_projr]   Úmid_channelsÚsqz_channelsÚ_layersÚ	__class__s                 €r#   rY   zMBConv.__init__D   st  ø€ ô 	‰ÑÔñ 	à ‘kÒ@ [°LÑ%@ˆÙÜ—I‘I˜k¨<ÀQÈqÐW[Ô\Ð]ˆDØ˜Š{ÜŸ™°¸6È1ÔMÐNÐQUÑU�ÜŸ™ tÐ,ˆD�IäŸ™›ˆDŒIä˜<¨/Ñ9Ó:ˆÜ˜<¨-Ñ7Ó8ˆáÜ$3Ð4HÈuÔ$UˆDÕ!ä$&§K¡K£MˆDÔ!ä“-ˆÙ(¨Ó5ˆ�
ÑÜ0ØØØØØØ-Ø!Øô	
ˆ�Ñô 1ØØØØØØ-Ø!ØØô

ˆ�Ñô ):¸,ÈÔac×ahÑahÔ(iˆÐ$Ñ%ÜŸI™I°,È\ÐghÐosÔtˆ�Ñä—m‘m GÓ,ˆ�r%   Úxc                 ón   — | j                  |«      }| j                  | j                  |«      «      }||z   S )zÍ
        Args:
            x (Tensor): Input tensor with expected layout of [B, C, H, W].
        Returns:
            Tensor: Output tensor with expected layout of [B, C, H / stride, W / stride].
        )r]   r`   rb   ©rc   ri   Úress      r#   ÚforwardzMBConv.forward�   s2   € ð �i‰i˜‹lˆØ×!Ñ! $§+¡+¨a£.Ó1ˆØ�Q‰wˆr%   )ç        )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r_   Úfloatr   r	   ÚModulerY   r
   rm   Ú__classcell__©rh   s   @r#   rB   rB   6   s˜   ø„ ñð, '*ñ;-àð;-ð ð;-ð ð	;-ð
 ð;-ð ð;-ð # 3¨¯	©	 >Ñ2ð;-ð ˜S "§)¡)˜^Ñ,ð;-ð $ð;-ð 
õ;-ðz	˜ð 	 F÷ 	r%   rB   c                   ód   ‡ — e Zd ZdZdedededdfˆ fd„Zdej                  fd„Zd	edefd
„Z	ˆ xZ
S )Ú$RelativePositionalMultiHeadAttentionzÔRelative Positional Multi-Head Attention.

    Args:
        feat_dim (int): Number of input features.
        head_dim (int): Number of features per head.
        max_seq_len (int): Maximum sequence length.
    Úfeat_dimÚhead_dimÚmax_seq_lenr   Nc                 ób  •— t         ‰| �  «        ||z  dk7  rt        d|› d|› �«      ‚||z  | _        || _        t        t        j                  |«      «      | _        || _	        t        j                  || j                  | j                  z  dz  «      | _        |dz  | _        t        j                  | j                  | j                  z  |«      | _        t        j                  j!                  t#        j$                  d| j                  z  dz
  d| j                  z  dz
  z  | j                  ft"        j&                  ¬«      «      | _        | j+                  d	t-        | j                  | j                  «      «       t"        j                  j.                  j1                  | j(                  d
¬«       y )Nr   z
feat_dim: z  must be divisible by head_dim: r(   g      à¿r    r!   )ÚdtypeÚrelative_position_indexç{®Gáz”?©Ústd)rX   rY   Ú
ValueErrorÚn_headsrz   r_   ÚmathÚsqrtÚsizer{   r	   ÚLinearÚto_qkvÚscale_factorÚmergeÚ	parameterÚ	Parameterr5   ÚemptyÚfloat32Úrelative_position_bias_tableÚregister_bufferr@   ÚinitÚtrunc_normal_)rc   ry   rz   r{   rh   s       €r#   rY   z-RelativePositionalMultiHeadAttention.__init__–   sU  ø€ ô 	‰ÑÔà�hÑ !Ò#Ü˜z¨(¨Ð3SÐT\ÐS]Ð^Ó_Ð_à 8Ñ+ˆŒØ ˆŒÜœŸ	™	 +Ó.Ó/ˆŒ	Ø&ˆÔä—i‘i ¨$¯,©,¸¿¹Ñ*FÈÑ*JÓKˆŒØ$ d™NˆÔä—Y‘Y˜tŸ}™}¨t¯|©|Ñ;¸XÓFˆŒ
Ü,.¯L©L×,BÑ,BÜ�K‰K˜!˜dŸi™i™-¨!Ñ+°°D·I±I±ÀÑ0AÑBÀDÇLÁLÐQÔY^×YfÑYfÔgó-
ˆÔ)ð 	×ÑÐ6Ô8TÐUY×U^ÑU^Ð`d×`iÑ`iÓ8jÔkä�‰�‰×#Ñ# D×$EÑ$EÈ4Ð#ÕPr%   c                 ó  — | j                   j                  d«      }| j                  |   j                  | j                  | j                  d«      }|j	                  ddd«      j                  «       }|j                  d«      S )Nr4   r    r   r!   )r~   Úviewr�   r{   r:   r;   Ú	unsqueeze)rc   Ú
bias_indexÚrelative_biass      r#   Úget_relative_positional_biaszARelativePositionalMultiHeadAttention.get_relative_positional_bias²   ss   € Ø×1Ñ1×6Ñ6°rÓ:ˆ
Ø×9Ñ9¸*ÑE×JÑJÈ4×K[ÑK[Ð]a×]mÑ]mÐoqÓrˆØ%×-Ñ-¨a°°AÓ6×AÑAÓCˆØ×&Ñ& qÓ)Ð)r%   ri   c                 óà  — |j                   \  }}}}| j                  | j                  }}| j                  |«      }t	        j
                  |dd¬«      \  }	}
}|	j                  |||||«      j                  ddddd«      }	|
j                  |||||«      j                  ddddd«      }
|j                  |||||«      j                  ddddd«      }|
| j                  z  }
t	        j                  d|	|
«      }| j                  «       }t        j                  ||z   d¬«      }t	        j                  d	||«      }|j                  ddddd«      j                  ||||«      }| j                  |«      }|S )
z»
        Args:
            x (Tensor): Input tensor with expected layout of [B, G, P, D].
        Returns:
            Tensor: Output tensor with expected layout of [B, G, P, D].
        r(   r4   )Údimr   r!   r    é   z!B G H I D, B G H J D -> B G H I Jz!B G H I J, B G H J D -> B G H I D)Úshaperƒ   rz   rˆ   r5   ÚchunkÚreshaper:   r‰   Úeinsumr˜   ÚFÚsoftmaxrŠ   )rc   ri   ÚBÚGÚPÚDÚHÚDHÚqkvÚqÚkÚvÚdot_prodÚpos_biasÚouts                  r#   rm   z,RelativePositionalMultiHeadAttention.forward¸   sX  € ð —W‘W‰
ˆˆ1ˆa�Ø—‘˜dŸm™mˆ2ˆà�k‰k˜!‹nˆÜ—+‘+˜c 1¨"Ô-‰ˆˆ1ˆaà�I‰I�a˜˜A˜q "Ó%×-Ñ-¨a°°A°q¸!Ó<ˆØ�I‰I�a˜˜A˜q "Ó%×-Ñ-¨a°°A°q¸!Ó<ˆØ�I‰I�a˜˜A˜q "Ó%×-Ñ-¨a°°A°q¸!Ó<ˆà�×!Ñ!Ñ!ˆÜ—<‘<Ð CÀQÈÓJˆØ×4Ñ4Ó6ˆä—9‘9˜X¨Ñ0°bÔ9ˆä�l‰lÐ>ÀÈ!ÓLˆØ�k‰k˜!˜Q  1 aÓ(×0Ñ0°°A°q¸!Ó<ˆà�j‰j˜‹oˆØˆ
r%   )ro   rp   rq   rr   r_   rY   r5   r
   r˜   rm   ru   rv   s   @r#   rx   rx   �   s[   ø„ ñðQàðQð ðQð ð	Qð
 
õQð8*¨e¯l©ló *ð˜ð  F÷ r%   rx   c                   óh   ‡ — e Zd ZdZdededdfˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )	ÚSwapAxeszPermute the axes of a tensor.ÚaÚbr   Nc                 ó>   •— t         ‰| �  «        || _        || _        y ©N)rX   rY   r±   r²   )rc   r±   r²   rh   s      €r#   rY   zSwapAxes.__init__Ù   s   ø€ Ü‰ÑÔØˆŒØˆ�r%   ri   c                 ó\   — t        j                  || j                  | j                  «      }|S r´   )r5   Úswapaxesr±   r²   rk   s      r#   rm   zSwapAxes.forwardÞ   s!   € Ü�n‰n˜Q §¡¨¯©Ó/ˆØˆ
r%   )
ro   rp   rq   rr   r_   rY   r5   r
   rm   ru   rv   s   @r#   r°   r°   Ö   s;   ø„ Ù'ð˜#ð  #ð ¨$õ ð
˜Ÿ™ð ¨%¯,©,÷ r%   r°   c                   ó8   ‡ — e Zd ZdZdˆ fd„Zdededefd„Zˆ xZS )ÚWindowPartitionzB
    Partition the input tensor into non-overlapping windows.
    r   c                 ó"   •— t         ‰| �  «        y r´   ©rX   rY   ©rc   rh   s    €r#   rY   zWindowPartition.__init__è   ó   ø€ Ü‰ÑÕr%   ri   Úpc                 óÐ   — |j                   \  }}}}|}|j                  ||||z  |||z  |«      }|j                  dddddd«      }|j                  |||z  ||z  z  ||z  |«      }|S )zï
        Args:
            x (Tensor): Input tensor with expected layout of [B, C, H, W].
            p (int): Number of partitions.
        Returns:
            Tensor: Output tensor with expected layout of [B, H/P, W/P, P*P, C].
        r   r    r›   r(   é   r!   ©rœ   rž   r:   )rc   ri   r½   r¢   ÚCr¦   ÚWr¤   s           r#   rm   zWindowPartition.forwardë   s|   € ð —W‘W‰
ˆˆ1ˆa�Øˆà�I‰I�a˜˜A ™F A q¨A¡v¨qÓ1ˆØ�I‰I�a˜˜A˜q ! QÓ'ˆà�I‰I�a˜!˜q™& Q¨!¡VÑ,¨a°!©e°QÓ7ˆØˆr%   ©r   N©	ro   rp   rq   rr   rY   r
   r_   rm   ru   rv   s   @r#   r¸   r¸   ã   s'   ø„ ñõð˜ð  Cð ¨F÷ r%   r¸   c            
       ó@   ‡ — e Zd ZdZd	ˆ fd„Zdededededef
d„Zˆ xZS )
ÚWindowDepartitionzo
    Departition the input tensor of non-overlapping windows into a feature volume of layout [B, C, H, W].
    r   c                 ó"   •— t         ‰| �  «        y r´   rº   r»   s    €r#   rY   zWindowDepartition.__init__  r¼   r%   ri   r½   Úh_partitionsÚw_partitionsc                 óÀ   — |j                   \  }}}}|}	||}}
|j                  ||
||	|	|«      }|j                  dddddd«      }|j                  |||
|	z  ||	z  «      }|S )ar  
        Args:
            x (Tensor): Input tensor with expected layout of [B, (H/P * W/P), P*P, C].
            p (int): Number of partitions.
            h_partitions (int): Number of vertical partitions.
            w_partitions (int): Number of horizontal partitions.
        Returns:
            Tensor: Output tensor with expected layout of [B, C, H, W].
        r   r¿   r!   r(   r    r›   rÀ   )rc   ri   r½   rÈ   rÉ   r¢   r£   ÚPPrÁ   r¤   ÚHPÚWPs               r#   rm   zWindowDepartition.forward  st   € ð —g‘g‰ˆˆ1ˆb�!ØˆØ˜|ˆBˆà�I‰I�a˜˜R  A qÓ)ˆà�I‰I�a˜˜A˜q ! QÓ'ˆà�I‰I�a˜˜B ™F B¨¡FÓ+ˆØˆr%   rÃ   rÄ   rv   s   @r#   rÆ   rÆ   ý   s6   ø„ ñõð˜ð  Cð °sð È#ð ÐRX÷ r%   rÆ   c                   óª   ‡ — e Zd ZdZdededededeeef   deded	ej                  f   d
ed	ej                  f   de
de
de
ddfˆ fd„Zdedefd„Zˆ xZS )ÚPartitionAttentionLayera®  
    Layer for partitioning the input tensor into non-overlapping windows and applying attention to each window.

    Args:
        in_channels (int): Number of input channels.
        head_dim (int): Dimension of each attention head.
        partition_size (int): Size of the partitions.
        partition_type (str): Type of partitioning to use. Can be either "grid" or "window".
        grid_size (Tuple[int, int]): Size of the grid to partition the input tensor into.
        mlp_ratio (int): Ratio of the  feature size expansion in the MLP layer.
        activation_layer (Callable[..., nn.Module]): Activation function to use.
        norm_layer (Callable[..., nn.Module]): Normalization function to use.
        attention_dropout (float): Dropout probability for the attention layer.
        mlp_dropout (float): Dropout probability for the MLP layer.
        p_stochastic_dropout (float): Probability of dropping out a partition.
    rC   rz   Úpartition_sizeÚpartition_typeÚ	grid_sizeÚ	mlp_ratiorG   .rH   Úattention_dropoutÚmlp_dropoutrI   r   Nc           	      ó„  •— t         ‰| �  «        ||z  | _        || _        |d   |z  | _        || _        || _        |dvrt        d«      ‚|dk(  r|| j                  c| _        | _	        n| j                  |c| _        | _	        t        «       | _        t        «       | _        |dk(  rt        dd«      nt        j                   «       | _        |dk(  rt        dd«      nt        j                   «       | _        t        j&                   ||«      t)        |||dz  «      t        j*                  |	«      «      | _        t        j&                  t        j.                  |«      t        j0                  |||z  «       |«       t        j0                  ||z  |«      t        j*                  |
«      «      | _        t5        |d	¬
«      | _        y )Nr   )ÚgridÚwindowz0partition_type must be either 'grid' or 'window'rØ   r×   éþÿÿÿéýÿÿÿr    rM   rN   )rX   rY   rƒ   rz   Ún_partitionsrÑ   rÒ   r‚   r½   Úgr¸   Úpartition_oprÆ   Údepartition_opr°   r	   r^   Úpartition_swapÚdepartition_swapr\   rx   ÚDropoutÚ
attn_layerÚ	LayerNormr‡   Ú	mlp_layerr   Ústochastic_dropout)rc   rC   rz   rÐ   rÑ   rÒ   rÓ   rG   rH   rÔ   rÕ   rI   rh   s               €r#   rY   z PartitionAttentionLayer.__init__-  s…  ø€ ô" 	‰ÑÔà" hÑ.ˆŒØ ˆŒØ% a™L¨NÑ:ˆÔØ,ˆÔØ"ˆŒàÐ!3Ñ3ÜÐOÓPÐPà˜XÒ%Ø+¨T×->Ñ->ˆNˆDŒF�D•Fà!×.Ñ.°ˆNˆDŒF�D”Fä+Ó-ˆÔÜ/Ó1ˆÔØ2@ÀFÒ2Jœh r¨2Ô.ÔPR×P[ÑP[ÓP]ˆÔØ4BÀfÒ4L¤¨¨RÔ 0ÔRT×R]ÑR]ÓR_ˆÔäŸ-™-Ù�{Ó#ô 1°¸hÈÐXYÑHYÓZÜ�J‰JÐ(Ó)ó
ˆŒô Ÿ™Ü�L‰L˜Ó%Ü�I‰I�k ;°Ñ#:Ó;ÙÓÜ�I‰I�k IÑ-¨{Ó;Ü�J‰J�{Ó#ó
ˆŒô #2Ð2FÈUÔ"SˆÕr%   ri   c                 óª  — | j                   d   | j                  z  | j                   d   | j                  z  }}t        j                  | j                   d   | j                  z  dk(  xr | j                   d   | j                  z  dk(  dj	                  | j                   | j                  «      «       | j                  || j                  «      }| j                  |«      }|| j                  | j                  |«      «      z   }|| j                  | j                  |«      «      z   }| j                  |«      }| j                  || j                  ||«      }|S )z»
        Args:
            x (Tensor): Input tensor with expected layout of [B, C, H, W].
        Returns:
            Tensor: Output tensor with expected layout of [B, C, H, W].
        r   r!   z[Grid size must be divisible by partition size. Got grid size of {} and partition size of {})rÒ   r½   r5   Ú_assertÚformatrÝ   rß   rå   râ   rä   rà   rÞ   )rc   ri   ÚghÚgws       r#   rm   zPartitionAttentionLayer.forwardg  s!  € ð —‘ Ñ" d§f¡fÑ,¨d¯n©n¸QÑ.?À4Ç6Á6Ñ.IˆBˆÜ�‰Ø�N‰N˜1Ñ §¡Ñ&¨!Ñ+ÒO°·±¸qÑ0AÀDÇFÁFÑ0JÈaÑ0OØi×pÑpØ—‘ §¡óô	
ð ×Ñ˜a §¡Ó(ˆØ×Ñ Ó"ˆØ�×'Ñ'¨¯©¸Ó(:Ó;Ñ;ˆØ�×'Ñ'¨¯©°qÓ(9Ó:Ñ:ˆØ×!Ñ! !Ó$ˆØ×Ñ  4§6¡6¨2¨rÓ2ˆàˆr%   )ro   rp   rq   rr   r_   ÚstrÚtupler   r	   rt   rs   rY   r
   rm   ru   rv   s   @r#   rÏ   rÏ     sÆ   ø„ ñð"8Tàð8Tð ð8Tð
 ð8Tð ð8Tð ˜˜c˜‘?ð8Tð ð8Tð # 3¨¯	©	 >Ñ2ð8Tð ˜S "§)¡)˜^Ñ,ð8Tð !ð8Tð ð8Tð $ð8Tð  
õ!8Tðt˜ð  F÷ r%   rÏ   c                   ó¶   ‡ — e Zd ZdZdededededededej                  f   d	edej                  f   d
edededededede	eef   ddfˆ fd„Z
dedefd„Zˆ xZS )ÚMaxVitLayera  
    MaxVit layer consisting of a MBConv layer followed by a PartitionAttentionLayer with `window` and a PartitionAttentionLayer with `grid`.

    Args:
        in_channels (int): Number of input channels.
        out_channels (int): Number of output channels.
        expansion_ratio (float): Expansion ratio in the bottleneck.
        squeeze_ratio (float): Squeeze ratio in the SE Layer.
        stride (int): Stride of the depthwise convolution.
        activation_layer (Callable[..., nn.Module]): Activation function.
        norm_layer (Callable[..., nn.Module]): Normalization function.
        head_dim (int): Dimension of the attention heads.
        mlp_ratio (int): Ratio of the MLP layer.
        mlp_dropout (float): Dropout probability for the MLP layer.
        attention_dropout (float): Dropout probability for the attention layer.
        p_stochastic_dropout (float): Probability of stochastic depth.
        partition_size (int): Size of the partitions.
        grid_size (Tuple[int, int]): Size of the input feature grid.
    rC   rD   rF   rE   r   rH   .rG   rz   rÓ   rÕ   rÔ   rI   rÐ   rÒ   r   Nc                 ó2  •— t         ‰| �  «        t        «       }t        ||||||||¬«      |d<   t	        |||d||	|t
        j                  ||
|¬«      |d<   t	        |||d||	|t
        j                  ||
|¬«      |d<   t        j                  |«      | _        y )N)rC   rD   rE   rF   r   rG   rH   rI   ÚMBconvrØ   )rC   rz   rÐ   rÑ   rÒ   rÓ   rG   rH   rÔ   rÕ   rI   Úwindow_attentionr×   Úgrid_attention)	rX   rY   r   rB   rÏ   r	   rã   r\   rb   )rc   rC   rD   rF   rE   r   rH   rG   rz   rÓ   rÕ   rÔ   rI   rÐ   rÒ   rb   rh   s                   €r#   rY   zMaxVitLayer.__init__˜  sÀ   ø€ ô* 	‰ÑÔä)›mˆô "Ø#Ø%Ø+Ø'ØØ-Ø!Ø!5ô	
ˆˆxÑô &=Ø$ØØ)Ø#ØØØ-Ü—|‘|Ø/Ø#Ø!5ô&
ˆÐ!Ñ"ô $;Ø$ØØ)Ø!ØØØ-Ü—|‘|Ø/Ø#Ø!5ô$
ˆÐÑ ô —m‘m FÓ+ˆ�r%   ri   c                 ó(   — | j                  |«      }|S ©z�
        Args:
            x (Tensor): Input tensor of shape (B, C, H, W).
        Returns:
            Tensor: Output tensor of shape (B, C, H, W).
        ©rb   )rc   ri   s     r#   rm   zMaxVitLayer.forwardÙ  s   € ð �K‰K˜‹NˆØˆr%   )ro   rp   rq   rr   r_   rs   r   r	   rt   rì   rY   r
   rm   ru   rv   s   @r#   rî   rî   ƒ  sÙ   ø„ ñð(?,ð ð?,ð ð	?,ð
 ð?,ð ð?,ð ð?,ð ˜S "§)¡)˜^Ñ,ð?,ð # 3¨¯	©	 >Ñ2ð?,ð ð?,ð ð?,ð ð?,ð !ð?,ð  $ð!?,ð$ ð%?,ð& ˜˜c˜‘?ð'?,ð( 
õ)?,ðB˜ð  F÷ r%   rî   c                   ó¼   ‡ — e Zd ZdZdedededededej                  f   dedej                  f   d	ed
edededede	eef   dede
e   ddfˆ fd„Zdedefd„Zˆ xZS )ÚMaxVitBlocka(  
    A MaxVit block consisting of `n_layers` MaxVit layers.

     Args:
        in_channels (int): Number of input channels.
        out_channels (int): Number of output channels.
        expansion_ratio (float): Expansion ratio in the bottleneck.
        squeeze_ratio (float): Squeeze ratio in the SE Layer.
        activation_layer (Callable[..., nn.Module]): Activation function.
        norm_layer (Callable[..., nn.Module]): Normalization function.
        head_dim (int): Dimension of the attention heads.
        mlp_ratio (int): Ratio of the MLP layer.
        mlp_dropout (float): Dropout probability for the MLP layer.
        attention_dropout (float): Dropout probability for the attention layer.
        p_stochastic_dropout (float): Probability of stochastic depth.
        partition_size (int): Size of the partitions.
        input_grid_size (Tuple[int, int]): Size of the input feature grid.
        n_layers (int): Number of layers in the block.
        p_stochastic (List[float]): List of probabilities for stochastic depth for each layer.
    rC   rD   rF   rE   rH   .rG   rz   rÓ   rÕ   rÔ   rÐ   Úinput_grid_sizeÚn_layersÚp_stochasticr   Nc                 óp  •— t         ‰| �  «        t        |«      |k(  st        d|› d|› d�«      ‚t	        j
                  «       | _        t        |ddd¬«      | _        t        |«      D ]L  \  }}|dk(  rdnd}| xj                  t        |dk(  r|n||||||||||	|
|| j                  |¬	«      gz  c_        ŒN y )
Nz'p_stochastic must have length n_layers=z, got p_stochastic=ú.r(   r    r!   rL   r   )rC   rD   rF   rE   r   rH   rG   rz   rÓ   rÕ   rÔ   rÐ   rÒ   rI   )rX   rY   Úlenr‚   r	   Ú
ModuleListrb   r$   rÒ   Ú	enumeraterî   )rc   rC   rD   rF   rE   rH   rG   rz   rÓ   rÕ   rÔ   rÐ   rø   rù   rú   Úidxr½   r   rh   s                     €r#   rY   zMaxVitBlock.__init__ú  sÓ   ø€ ô, 	‰ÑÔÜ�<Ó  HÒ,ÜÐFÀxÀjÐPcÐdpÐcqÐqrÐsÓtÐtä—m‘m“oˆŒä/°ÈQÐWXÐbcÔdˆŒä Ó-ò 	‰FˆC�Ø š(‘Q¨ˆFØ�KŠKÜØ/2°aªx¡¸\Ø!-Ø"/Ø$3Ø!Ø)Ø%5Ø%Ø'Ø +Ø&7Ø#1Ø"Ÿn™nØ)*ôðñ ŽKñ	r%   ri   c                 ó8   — | j                   D ]
  } ||«      }Œ |S rô   rõ   )rc   ri   Úlayers      r#   rm   zMaxVitBlock.forward-  s%   € ð —[‘[ò 	ˆEÙ�a“‰Að	àˆr%   )ro   rp   rq   rr   r_   rs   r   r	   rt   rì   ÚlistrY   r
   rm   ru   rv   s   @r#   r÷   r÷   ä  sÝ   ø„ ñð*1ð ð1ð ð	1ð
 ð1ð ð1ð ˜S "§)¡)˜^Ñ,ð1ð # 3¨¯	©	 >Ñ2ð1ð ð1ð ð1ð ð1ð !ð1ð  ð!1ð" ˜s C˜x™ð#1ð& ð'1ð( ˜5‘kð)1ð* 
õ+1ðf	˜ð 	 F÷ 	r%   r÷   c            !       óø   ‡ — e Zd ZdZdej
                  ddddddfdeeef   ded	ed
ee   dee   dede	de
edej                  f      dedej                  f   de	de	dede	de	deddf ˆ fd„Zdedefd„Zd„ Zˆ xZS )r   ay  
    Implements MaxVit Transformer from the `MaxViT: Multi-Axis Vision Transformer <https://arxiv.org/abs/2204.01697>`_ paper.
    Args:
        input_size (Tuple[int, int]): Size of the input image.
        stem_channels (int): Number of channels in the stem.
        partition_size (int): Size of the partitions.
        block_channels (List[int]): Number of channels in each block.
        block_layers (List[int]): Number of layers in each block.
        stochastic_depth_prob (float): Probability of stochastic depth. Expands to a list of probabilities for each layer that scales linearly to the specified value.
        squeeze_ratio (float): Squeeze ratio in the SE Layer. Default: 0.25.
        expansion_ratio (float): Expansion ratio in the MBConv bottleneck. Default: 4.
        norm_layer (Callable[..., nn.Module]): Normalization function. Default: None (setting to None will produce a `BatchNorm2d(eps=1e-3, momentum=0.01)`).
        activation_layer (Callable[..., nn.Module]): Activation function Default: nn.GELU.
        head_dim (int): Dimension of the attention heads.
        mlp_ratio (int): Expansion ratio of the MLP layer. Default: 4.
        mlp_dropout (float): Dropout probability for the MLP layer. Default: 0.0.
        attention_dropout (float): Dropout probability for the attention layer. Default: 0.0.
        num_classes (int): Number of classes. Default: 1000.
    Ng      Ð?r›   rn   iè  r   Ústem_channelsrÐ   Úblock_channelsÚblock_layersrz   Ústochastic_depth_probrH   .rG   rF   rE   rÓ   rÕ   rÔ   Únum_classesr   c                 ó¦  •— t         ‰| �  «        t        | «       d}|€t        t        j
                  dd¬«      }t        |t        |«      «      }t        |«      D ]3  \  }}|d   |z  dk7  s|d   |z  dk7  sŒt        d|› d|› d	|› d
|› d�	«      ‚ t	        j                  t        ||dd||	dd ¬«      t        ||ddd d d¬«      «      | _        t        |ddd¬«      }|| _        t	        j                  «       | _        |g|d d z   }|}t#        j$                  d|t'        |«      «      j)                  «       }d}t+        |||«      D ]\  \  }}}| j                   j-                  t/        |||
|||	|||||||||||z    ¬«      «       | j                   d   j0                  }||z  }Œ^ t	        j                  t	        j2                  d«      t	        j4                  «       t	        j6                  |d   «      t	        j8                  |d   |d   «      t	        j:                  «       t	        j8                  |d   |d¬«      «      | _        | j?                  «        y )Nr(   gü©ñÒMbP?g{®Gáz„?)ÚepsÚmomentumr   r!   zInput size z
 of block z$ is not divisible by partition size zx. Consider changing the partition size or the input size.
Current configuration yields the following block input sizes: rü   r    F)r   rH   rG   rK   rQ   T)r   rH   rG   rK   rL   r4   )rC   rD   rF   rE   rH   rG   rz   rÓ   rÕ   rÔ   rÐ   rø   rù   rú   )rK   ) rX   rY   r   r   r	   ÚBatchNorm2dr.   rý   rÿ   r‚   r\   r   Ústemr$   rÐ   rþ   ÚblocksÚnpÚlinspacer<   ÚtolistÚzipr*   r÷   rÒ   ÚAdaptiveAvgPool2dÚFlattenrã   r‡   ÚTanhÚ
classifierÚ_init_weights)rc   r   r  rÐ   r  r  rz   r  rH   rG   rF   rE   rÓ   rÕ   rÔ   r	  Úinput_channelsÚblock_input_sizesr   Úblock_input_sizerC   rD   rú   Úp_idxÚ
in_channelÚout_channelÚ
num_layersrh   s                              €r#   rY   zMaxVit.__init__N  s‹  ø€ ô: 	‰ÑÔÜ˜DÔ!àˆð ÐÜ ¤§¡°TÀDÔIˆJô
 5°ZÄÀ^ÓATÓUÐÜ%.Ð/@Ó%Aò 	Ñ!ˆCÐ!Ø Ñ" ^Ñ3°qÒ8Ð<LÈQÑ<OÐR`Ñ<`ÐdeÓ<eÜ Ø!Ð"2Ð!3°:¸c¸UÐBfÐguÐfvð wUàUfÐTgÐghðjóð ð	ô —M‘MÜ ØØØØØ%Ø!1ØØô	ô !Ø˜}¨a¸ÀdÐ]aÐhlôó
ˆŒ	ô" ,¨JÀAÈaÐYZÔ[ˆ
Ø,ˆÔô —m‘m“oˆŒØ$�o¨°s¸Ð(;Ñ;ˆØ%ˆô
 —{‘{ 1Ð&;¼SÀÓ=NÓO×VÑVÓXˆàˆÜ36°{ÀLÐR^Ó3_ò 	 Ñ/ˆJ˜ ZØ�K‰K×ÑÜØ *Ø!,Ø"/Ø$3Ø)Ø%5Ø%Ø'Ø +Ø&7Ø#1Ø$.Ø'Ø!-¨e°e¸jÑ6HÐ!Iôôð$ Ÿ™ R™×2Ñ2ˆJØ�ZÑ‰Eð)	 ô0 Ÿ-™-Ü× Ñ  Ó#Ü�J‰J‹LÜ�L‰L˜¨Ñ+Ó,Ü�I‰I�n RÑ(¨.¸Ñ*<Ó=Ü�G‰G‹IÜ�I‰I�n RÑ(¨+¸EÔBó
ˆŒð 	×ÑÕr%   ri   c                 ó|   — | j                  |«      }| j                  D ]
  } ||«      }Œ | j                  |«      }|S r´   )r  r  r  )rc   ri   Úblocks      r#   rm   zMaxVit.forwardÄ  s>   € Ø�I‰I�a‹LˆØ—[‘[ò 	ˆEÙ�a“‰Að	à�O‰O˜AÓˆØˆr%   c                 ó  — | j                  «       D �]l  }t        |t        j                  «      rbt        j                  j                  |j                  d¬«       |j                  €ŒVt        j                  j                  |j                  «       Œ€t        |t        j                  «      rUt        j                  j                  |j                  d«       t        j                  j                  |j                  d«       Œït        |t        j                  «      s�Œt        j                  j                  |j                  d¬«       |j                  €�ŒDt        j                  j                  |j                  «       �Œo y )Nr   r€   r!   r   )ÚmodulesÚ
isinstancer	   rZ   r‘   Únormal_ÚweightrK   Úzeros_r  Ú	constant_r‡   )rc   Úms     r#   r  zMaxVit._init_weightsË  sÝ   € Ø—‘“ó 	+ˆAÜ˜!œRŸY™YÔ'Ü—‘—‘ §¡¨d�Ô3Ø—6‘6Ñ%Ü—G‘G—N‘N 1§6¡6Õ*Ü˜AœrŸ~™~Ô.Ü—‘×!Ñ! !§(¡(¨AÔ.Ü—‘×!Ñ! !§&¡&¨!Õ,Ü˜AœrŸy™yÖ)Ü—‘—‘ §¡¨d�Ô3Ø—6‘6Ò%Ü—G‘G—N‘N 1§6¡6Ö*ñ	+r%   )ro   rp   rq   rr   r	   ÚGELUrì   r_   r  rs   r   r   rt   rY   r
   rm   r  ru   rv   s   @r#   r   r   9  s%  ø„ ñðJ :>Ø57·W±Wà#Ø!"àØ Ø#&àñ7tð ˜#˜s˜(‘Oðtð
 ðtð ðtð ˜S™	ðtð ˜3‘iðtð ðtð  %ðtð" ˜X c¨2¯9©9 nÑ5Ñ6ð#tð$ # 3¨¯	©	 >Ñ2ð%tð( ð)tð* ð+tð. ð/tð0 ð1tð2 !ð3tð6 ð7tð8 
õ9tðl˜ð  Fó ö+r%   r   r  r  r  r  rÐ   rz   ÚweightsÚprogressÚkwargsc                 ód  — |�dt        |dt        |j                  d   «      «       |j                  d   d   |j                  d   d   k(  sJ ‚t        |d|j                  d   «       |j                  dd«      }	t	        d| ||||||	dœ|¤Ž}
|�"|
j                  |j                  |d	¬
«      «       |
S )Nr	  Ú
categoriesÚmin_sizer   r!   r   ©éà   r2  )r  r  r  r  rz   rÐ   r   T)r,  Ú
check_hashr"   )r   rý   ÚmetaÚpopr   Úload_state_dictÚget_state_dict)r  r  r  r  rÐ   rz   r+  r,  r-  r   Úmodels              r#   Ú_maxvitr9  Ú  sÍ   € ð$ ÐÜ˜f m´S¸¿¹ÀlÑ9SÓ5TÔUØ�|‰|˜JÑ'¨Ñ*¨g¯l©l¸:Ñ.FÀqÑ.IÒIÐIÐIÜ˜f l°G·L±LÀÑ4LÔMà—‘˜L¨*Ó5€Jäð 	Ø#Ø%Ø!Ø3ØØ%Øñ	ð ñ	€Eð ÐØ×Ñ˜g×4Ñ4¸hÐSWÐ4ÓXÔYà€Lr%   c                   ój   — e Zd Z ed eeddej                  ¬«      edddddd	d
œiddddœ¬«      Z	e	Z
y)r   z9https://download.pytorch.org/models/maxvit_t-bc5ab103.pthr2  )Ú	crop_sizeÚresize_sizeÚinterpolationiÈË×r1  zLhttps://github.com/pytorch/vision/tree/main/references/classification#maxvitzImageNet-1KgÍÌÌÌÌìT@g‘í|?5.X@)zacc@1zacc@5g¬Zd;@gð§ÆK7±]@z½These weights reproduce closely the results of the paper using a similar training recipe.
            They were trained with a BatchNorm2D momentum of 0.99 instead of the more correct 0.01.)r/  Ú
num_paramsr0  ÚrecipeÚ_metricsÚ_opsÚ
_file_sizeÚ_docs)ÚurlÚ
transformsr4  N)ro   rp   rq   r   r   r   r   ÚBICUBICr   ÚIMAGENET1K_V1ÚDEFAULTr"   r%   r#   r   r     sb   „ ÙàGÙØ¨3¸CÐO`×OhÑOhô
ð /Ø"Ø"ØdàØ#Ø#ñ ðð Ø!ðgñ
ô€Mð. �Gr%   r   Ú
pretrained)r+  T)r+  r,  c                 ó\   — t         j                  | «      } t        ddg d¢g d¢ddd| |dœ|¤ŽS )	aŠ  
    Constructs a maxvit_t architecture from
    `MaxViT: Multi-Axis Vision Transformer <https://arxiv.org/abs/2204.01697>`_.

    Args:
        weights (:class:`~torchvision.models.MaxVit_T_Weights`, optional): The
            pretrained weights to use. See
            :class:`~torchvision.models.MaxVit_T_Weights` below for
            more details, and possible values. By default, no pre-trained
            weights are used.
        progress (bool, optional): If True, displays a progress bar of the
            download to stderr. Default is True.
        **kwargs: parameters passed to the ``torchvision.models.maxvit.MaxVit``
            base class. Please refer to the `source code
            <https://github.com/pytorch/vision/blob/main/torchvision/models/maxvit.py>`_
            for more details about this class.

    .. autoclass:: torchvision.models.MaxVit_T_Weights
        :members:
    é@   )rK  é€   é   i   )r    r    r¿   r    é    gš™™™™™É?é   )r  r  r  rz   r  rÐ   r+  r,  r"   )r   Úverifyr9  )r+  r,  r-  s      r#   r   r     sH   € ô. ×%Ñ% gÓ.€Gäð 
ØÚ*Ú!ØØ!ØØØñ
ð ñ
ð 
r%   )NF)=r„   Úcollectionsr   Úcollections.abcr   Ú	functoolsr   Útypingr   r   r   Únumpyr  r5   Útorch.nn.functionalr	   Ú
functionalr    r
   Útorchvision.models._apir   r   r   Útorchvision.models._metar   Útorchvision.models._utilsr   r   Útorchvision.ops.miscr   r   Ú torchvision.ops.stochastic_depthr   Útorchvision.transforms._presetsr   r   Útorchvision.utilsr   Ú__all__rì   r_   r$   r  r.   r@   rt   rB   rx   r°   r¸   rÆ   rÏ   rî   r÷   r   rs   Úboolr9  r   rG  r   r"   r%   r#   ú<module>ra     s[  ðÛ Ý #Ý $Ý ß *Ñ *ã Û ß Ð ß ß HÑ HÝ 9ß Tß HÝ <ß RÝ 1ò€ð u¨S°#¨X¡ð ÀSð ÐRUð Ð`cð ÐhmÐnqÐsvÐnvÑhwó ð¨¨s°C¨x©ð ÀCð ÈDÐQVÐWZÐ\_ÐW_ÑQ`ÑLaó ð#¨ð #°Sð #¸U¿\¹\ó #ôTˆR�Y‰Yô TônF¨2¯9©9ô FôR
ˆr�y‰yô 
ô�b—i‘iô ô4˜Ÿ	™	ô ô<e˜bŸi™iô eôP^�"—)‘)ô ^ôBR�"—)‘)ô Rôj^+ˆR�Y‰Yô ^+ðZ &*Øñ'àð'ð ˜‘Ið	'ð
 �s‘)ð'ð !ð'ð ð'ð ð'ð �kÑ"ð'ð ð'ð ð'ð  ó!'ôT�{ô ñ6 ÓÙ ,Ð0@×0NÑ0NÐ!OÔPØ6:ÈTò !˜Ð"2Ñ3ð !Àdð !Ð]`ð !Ðekò !ó Qó ñ!r%   