ó
    Eñin}  ã                   ó¶  • S SK r S SKJr  S SKJr  S SKJr  S SKJrJ	r	J
r
  S SKrS SKrS SKJs  Jr  S SKJrJr  S SKJrJrJr  S SKJr  S S	KJrJr  S S
KJrJr  S SKJ r   S SK!J"r"J#r#  S SK$J%r%  / SQr&S\'\(\(4   S\(S\(S\(S\'\(\(4   4
S jr)S\'\(\(4   S\(S\*\'\(\(4      4S jr+S\(S\(S\R$                  4S jr, " S S\RZ                  5      r. " S S\RZ                  5      r/ " S S\RZ                  5      r0 " S  S!\RZ                  5      r1 " S" S#\RZ                  5      r2 " S$ S%\RZ                  5      r3 " S& S'\RZ                  5      r4 " S( S)\RZ                  5      r5 " S* S+\RZ                  5      r6  S=S,\(S-\*\(   S.\*\(   S/\7S0\(S1\(S2\
\   S3\8S4\S\64S5 jjr9 " S6 S7\5      r:\" 5       \" S8\:Rv                  4S99SS:S;.S2\
\:   S3\8S4\S\64S< jj5       5       r<g)>é    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   • U S   U-
  SU-  -   U-  S-   U S   U-
  SU-  -   U-  S-   4$ )Nr   é   é   © )r   r   r   r   s       ÚV/home/mande/repo/quber/.venv/lib/python3.13/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                 óˆ   • / n[        U SSS5      n[        U5       H"  n[        USSS5      nUR                  U5        M$     U$ )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.   !   sL   € à€FÜ.¨z¸1¸aÀÓCÐÜ�8Ž_ˆÜ2Ð3DÀaÈÈAÓNÐØ�‰Ð'Ö(ñ ð €Mr%   ÚheightÚwidthc                 óü  • [         R                  " [         R                  " [         R                  " U 5      [         R                  " U5      /SS95      n[         R                  " US5      nUS S 2S S 2S 4   US S 2S S S 24   -
  nUR                  SSS5      R                  5       nUS S 2S S 2S4==   U S-
  -  ss'   US S 2S S 2S4==   US-
  -  ss'   US S 2S S 2S4==   SU-  S-
  -  ss'   UR                  S5      $ )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                   ó¨   ^ • \ rS rSrSr SS\S\S\S\S\S\S	\R                  4   S
\S	\R                  4   S\SS4U 4S jjjr
S\S\4S jrSrU =r$ )ÚMBConvé6   a  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U ]  5         U   US:g  =(       d    X:g  n	U	(       aQ  [        R                  " XSSSS9/n
US:X  a  [        R                  " SUSS9/U
-   n
[        R
                  " U
6 U l        O[        R                  " 5       U l        [        X#-  5      n[        X$-  5      nU(       a  [        USS9U l
        O[        R                  " 5       U l
        [        5       nU" U5      US	'   [        UUSSS
UUS S9US'   [        UUSUSUUUS S9	US'   [        X¼[        R                  S9US'   [        R                  " X²SSS9US'   [        R
                  " U5      U l        g )Nr!   T)r   r   Úbiasr    r(   ©r   r   r   Úrow©ÚmodeÚpre_normr   )r   r   r   rH   rI   ÚinplaceÚconv_a)r   r   r   rH   rI   ÚgroupsrR   Úconv_b)Ú
activationÚsqueeze_excitation)rD   rE   r   rL   Úconv_c)ÚsuperÚ__init__r	   ÚConv2dÚ	AvgPool2dÚ
SequentialÚprojÚIdentityÚintr   Ústochastic_depthr   r   r   ÚSiLUÚlayers)ÚselfrD   rE   rF   rG   r   rH   rI   rJ   Úshould_projr^   Úmid_channelsÚsqz_channelsÚ_layersÚ	__class__s                 €r#   rZ   ÚMBConv.__init__D   sh  ø€ ô 	‰ÑÔñ 	à ‘k×@ [Ñ%@ˆÞÜ—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                 ól   • U R                  U5      nU R                  U R                  U5      5      nX!-   $ )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^   ra   rc   ©rd   rk   Úress      r#   ÚforwardÚMBConv.forward�   s0   € ð �i‰i˜‹lˆØ×!Ñ! $§+¡+¨a£.Ó1ˆØ‰wˆr%   )rc   r^   ra   )ç        )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r`   Úfloatr   r	   ÚModulerZ   r
   ro   Ú__static_attributes__Ú__classcell__©ri   s   @r#   rB   rB   6   s¢   ø† ñð, '*ñ;-àð;-ð ð;-ð ð	;-ð
 ð;-ð ð;-ð # 3¨¯	©	 >Ñ2ð;-ð ˜S "§)¡)˜^Ñ,ð;-ð $ð;-ð 
÷;-ð ;-ðz	˜ð 	 F÷ 	ò 	r%   rB   c                   ót   ^ • \ rS rSrSrS\S\S\SS4U 4S jjrS\R                  4S	 jr	S
\S\4S jr
SrU =r$ )Ú$RelativePositionalMultiHeadAttentioné�   zÀ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                 óV  >• [         TU ]  5         X-  S:w  a  [        SU SU 35      eX-  U l        X l        [        [        R                  " U5      5      U l        X0l	        [        R                  " XR                  U R                  -  S-  5      U l        US-  U l        [        R                  " U R                  U R                  -  U5      U l        [        R                  R!                  ["        R$                  " SU R                  -  S-
  SU R                  -  S-
  -  U R                  4["        R&                  S95      U l        U R+                  S	[-        U R                  U R                  5      5        ["        R                  R.                  R1                  U R(                  S
S9  g )Nr   z
feat_dim: z  must be divisible by head_dim: r(   g      à¿r    r!   )ÚdtypeÚrelative_position_indexç{®Gáz”?©Ústd)rY   rZ   Ú
ValueErrorÚn_headsr€   r`   ÚmathÚsqrtÚsizer�   r	   ÚLinearÚto_qkvÚscale_factorÚmergeÚ	parameterÚ	Parameterr5   ÚemptyÚfloat32Úrelative_position_bias_tableÚregister_bufferr@   ÚinitÚtrunc_normal_)rd   r   r€   r�   ri   s       €r#   rZ   Ú-RelativePositionalMultiHeadAttention.__init__–   sK  ø€ ô 	‰ÑÔàÑ !Ó#Ü˜z¨(¨Ð3SÐT\ÐS]Ð^Ó_Ð_àÑ+ˆŒØ ŒÜœŸ	š	 +Ó.Ó/ˆŒ	Ø&Ôä—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                 ó  • U R                   R                  S5      nU R                  U   R                  U R                  U R                  S5      nUR	                  SSS5      R                  5       nUR                  S5      $ )Nr4   r    r   r!   )r„   Úviewr•   r�   r:   r;   Ú	unsqueeze)rd   Ú
bias_indexÚrelative_biass      r#   Úget_relative_positional_biasÚARelativePositionalMultiHeadAttention.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%   rk   c                 ó¼  • UR                   u  p#pEU R                  U R                  pvU R                  U5      n[        R
                  " USSS9u  pšnU	R                  X#XFU5      R                  SSSSS5      n	U
R                  X#XFU5      R                  SSSSS5      n
UR                  X#XFU5      R                  SSSSS5      nX R                  -  n
[        R                  " SXš5      nU R                  5       n[        R                  " XÍ-   SS9n[        R                  " S	XË5      nUR                  SSSSS5      R                  X#XE5      nU R                  U5      nU$ )
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‰   r€   rŽ   r5   ÚchunkÚreshaper:   r�   ÚeinsumrŸ   ÚFÚsoftmaxr�   )rd   rk   ÚBÚGÚPÚDÚHÚDHÚqkvÚqÚkÚvÚdot_prodÚpos_biasÚouts                  r#   ro   Ú,RelativePositionalMultiHeadAttention.forward¸   s8  € ð —W‘W‰
ˆˆaØ—‘˜dŸm™mˆ2à�k‰k˜!‹nˆÜ—+’+˜c 1¨"Ñ-‰ˆˆaà�I‰I�a˜A "Ó%×-Ñ-¨a°°A°q¸!Ó<ˆØ�I‰I�a˜A "Ó%×-Ñ-¨a°°A°q¸!Ó<ˆØ�I‰I�a˜A "Ó%×-Ñ-¨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°°qÓ<ˆà�j‰j˜‹oˆØˆ
r%   )r€   r�   r�   r‰   r•   r�   rŒ   rŽ   )rr   rs   rt   ru   rv   r`   rZ   r5   r
   rŸ   ro   ry   rz   r{   s   @r#   r}   r}   �   s`   ø† ñðQàðQð ðQð ð	Qð
 
÷Qð8*¨e¯l©lô *ð˜ð  F÷ ò r%   r}   c                   óv   ^ • \ rS rSrSrS\S\SS4U 4S jjrS\R                  S\R                  4S	 jr	S
r
U =r$ )ÚSwapAxeséÖ   zPermute the axes of a tensor.ÚaÚbr   Nc                 ó:   >• [         TU ]  5         Xl        X l        g ©N)rY   rZ   r»   r¼   )rd   r»   r¼   ri   s      €r#   rZ   ÚSwapAxes.__init__Ù   s   ø€ Ü‰ÑÔØŒØ�r%   rk   c                 ó\   • [         R                  " XR                  U R                  5      nU$ r¾   )r5   Úswapaxesr»   r¼   rm   s      r#   ro   ÚSwapAxes.forwardÞ   s   € Ü�nŠn˜Q§¡¨¯©Ó/ˆØˆ
r%   )r»   r¼   )rr   rs   rt   ru   rv   r`   rZ   r5   r
   ro   ry   rz   r{   s   @r#   r¹   r¹   Ö   s@   ø† Ù'ð˜#ð  #ð ¨$÷ ð
˜Ÿ™ð ¨%¯,©,÷ ò r%   r¹   c                   óF   ^ • \ rS rSrSrS	U 4S jjrS\S\S\4S jrSr	U =r
$ )
ÚWindowPartitionéã   z:
Partition the input tensor into non-overlapping windows.
r   c                 ó"   >• [         TU ]  5         g r¾   ©rY   rZ   ©rd   ri   s    €r#   rZ   ÚWindowPartition.__init__è   ó   ø€ Ü‰ÑÕr%   rk   Úpc                 óÀ   • UR                   u  p4pVUnUR                  X4XW-  XvU-  U5      nUR                  SSSSSS5      nUR                  X5U-  Xg-  -  Xw-  U5      nU$ )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:   )rd   rk   rË   rª   ÚCr®   ÚWr¬   s           r#   ro   ÚWindowPartition.forwardë   sl   € ð —W‘W‰
ˆˆaØˆà�I‰I�a˜A™F A¨A¡v¨qÓ1ˆØ�I‰I�a˜˜A˜q ! QÓ'ˆà�I‰I�a˜q™& Q¡VÑ,¨a©e°QÓ7ˆØˆr%   r"   ©r   N©rr   rs   rt   ru   rv   rZ   r
   r`   ro   ry   rz   r{   s   @r#   rÄ   rÄ   ã   s,   ø† ñ÷ð˜ð  Cð ¨F÷ ò r%   rÄ   c            
       óN   ^ • \ rS rSrSrSU 4S jjrS\S\S\S\S\4
S	 jrS
r	U =r
$ )ÚWindowDepartitionéý   zg
Departition the input tensor of non-overlapping windows into a feature volume of layout [B, C, H, W].
r   c                 ó"   >• [         TU ]  5         g r¾   rÇ   rÈ   s    €r#   rZ   ÚWindowDepartition.__init__  rÊ   r%   rk   rË   Úh_partitionsÚw_partitionsc                 ó¬   • UR                   u  pVpxUn	X4pºUR                  XZX¹X˜5      nUR                  SSSSSS5      nUR                  XXX©-  X¹-  5      nU$ )a2  
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Î   )rd   rk   rË   rÙ   rÚ   rª   r«   ÚPPrÏ   r¬   ÚHPÚWPs               r#   ro   ÚWindowDepartition.forward  s`   € ð —g‘g‰ˆˆbØˆØˆBà�I‰I�a˜R AÓ)ˆà�I‰I�a˜˜A˜q ! QÓ'ˆà�I‰I�a˜B™F B¡FÓ+ˆØˆr%   r"   rÒ   rÓ   r{   s   @r#   rÕ   rÕ   ý   s;   ø† ñ÷ð˜ð  Cð °sð È#ð ÐRX÷ ò r%   rÕ   c                   ó¸   ^ • \ rS rSrSrS\S\S\S\S\\\4   S\S	\S
\	R                  4   S\S
\	R                  4   S\S\S\SS4U 4S jjrS\S\4S jrSrU =r$ )ÚPartitionAttentionLayeri  av  
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.
rD   r€   Úpartition_sizeÚpartition_typeÚ	grid_sizeÚ	mlp_ratiorH   .rI   Úattention_dropoutÚmlp_dropoutrJ   r   Nc           	      óŠ  >• [         TU ]  5         X-  U l        X l        US   U-  U l        X@l        XPl        US;  a  [        S5      eUS:X  a  X0R                  sU l        U l	        OU R                  UsU l        U l	        [        5       U l        [        5       U l        US:X  a  [        SS5      O[        R                   " 5       U l        US:X  a  [        SS5      O[        R                   " 5       U l        [        R&                  " U" U5      [)        XUS-  5      [        R*                  " U	5      5      U l        [        R&                  " [        R.                  " U5      [        R0                  " XU-  5      U" 5       [        R0                  " X-  U5      [        R*                  " U
5      5      U l        [5        US	S
9U l        g )Nr   )ÚgridÚwindowz0partition_type must be either 'grid' or 'window'rê   ré   éþÿÿÿéýÿÿÿr    rN   rO   )rY   rZ   r‰   r€   Ún_partitionsrã   rä   rˆ   rË   ÚgrÄ   Úpartition_oprÕ   Údepartition_opr¹   r	   r_   Úpartition_swapÚdepartition_swapr]   r}   ÚDropoutÚ
attn_layerÚ	LayerNormr�   Ú	mlp_layerr   Ústochastic_dropout)rd   rD   r€   râ   rã   rä   rå   rH   rI   ræ   rç   rJ   ri   s               €r#   rZ   Ú PartitionAttentionLayer.__init__-  su  ø€ ô" 	‰ÑÔà"Ñ.ˆŒØ ŒØ% a™L¨NÑ:ˆÔØ,ÔØ"ŒàÐ!3Ó3ÜÐOÓPÐPà˜XÓ%Ø+×->Ñ->ˆ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°ÈÐXYÑHYÓZÜ�JŠJÐ(Ó)ó
ˆŒô ŸšÜ�LŠL˜Ó%Ü�IŠI�k°Ñ#:Ó;ÙÓÜ�IŠI�kÑ-¨{Ó;Ü�JŠJ�{Ó#ó
ˆŒô #2Ð2FÈUÑ"SˆÕr%   rk   c                 óª  • U R                   S   U R                  -  U R                   S   U R                  -  p2[        R                  " U R                   S   U R                  -  S:H  =(       a    U R                   S   U R                  -  S:H  SR	                  U R                   U R                  5      5        U R                  XR                  5      nU R                  U5      nXR                  U R                  U5      5      -   nXR                  U R                  U5      5      -   nU R                  U5      nU R                  XR                  X#5      nU$ )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ð   )rd   rk   ÚghÚgws       r#   ro   Ú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Ó:Ñ:ˆØ×!Ñ! !Ó$ˆØ×Ñ §6¡6¨2Ó2ˆàˆr%   )rô   rð   rò   rî   rä   r€   rö   r‰   rí   rË   rï   rñ   rã   r÷   )rr   rs   rt   ru   rv   r`   ÚstrÚtupler   r	   rx   rw   rZ   r
   ro   ry   rz   r{   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                   óÄ   ^ • \ rS rSrSrS\S\S\S\S\S\S	\R                  4   S
\S	\R                  4   S\S\S\S\S\S\S\
\\4   SS4U 4S jjrS\S\4S jrSrU =r$ )ÚMaxVitLayeriƒ  aÔ  
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.
rD   rE   rG   rF   r   rI   .rH   r€   rå   rç   ræ   rJ   râ   rä   r   Nc                 ó"  >• [         TU ]  5         [        5       n[        UUUUUUUUS9US'   [	        UUUSUU	U[
        R                  UU
US9US'   [	        UUUSUU	U[
        R                  UU
US9US'   [
        R                  " U5      U l        g )N)rD   rE   rF   rG   r   rH   rI   rJ   ÚMBconvrê   )rD   r€   râ   rã   rä   rå   rH   rI   ræ   rç   rJ   Úwindow_attentionré   Úgrid_attention)	rY   rZ   r   rB   rá   r	   rõ   r]   rc   )rd   rD   rE   rG   rF   r   rI   rH   r€   rå   rç   ræ   rJ   râ   rä   rc   ri   s                   €r#   rZ   ÚMaxVitLayer.__init__˜  sÀ   ø€ ô* 	‰ÑÔä)›mˆô "Ø#Ø%Ø+Ø'ØØ-Ø!Ø!5ñ	
ˆˆxÑô &=Ø$ØØ)Ø#ØØØ-Ü—|‘|Ø/Ø#Ø!5ñ&
ˆÐ!Ñ"ô $;Ø$ØØ)Ø!ØØØ-Ü—|‘|Ø/Ø#Ø!5ñ$
ˆÐÑ ô —m’m FÓ+ˆ�r%   rk   c                 ó(   • U R                  U5      nU$ ©zu
Args:
    x (Tensor): Input tensor of shape (B, C, H, W).
Returns:
    Tensor: Output tensor of shape (B, C, H, W).
©rc   )rd   rk   s     r#   ro   ÚMaxVitLayer.forwardÙ  s   € ð �K‰K˜‹NˆØˆr%   r
  )rr   rs   rt   ru   rv   r`   rw   r   r	   rx   r   rZ   r
   ro   ry   rz   r{   s   @r#   r  r  ƒ  sÞ   ø† ñð(?,ð ð?,ð ð	?,ð
 ð?,ð ð?,ð ð?,ð ˜S "§)¡)˜^Ñ,ð?,ð # 3¨¯	©	 >Ñ2ð?,ð ð?,ð ð?,ð ð?,ð !ð?,ð  $ð!?,ð$ ð%?,ð& ˜˜c˜‘?ð'?,ð( 
÷)?,ðB˜ð  F÷ ò r%   r  c                   óÊ   ^ • \ rS rSrSrS\S\S\S\S\S\R                  4   S	\S\R                  4   S
\S\S\S\S\S\
\\4   S\S\\   SS4U 4S jjrS\S\4S jrSrU =r$ )ÚMaxVitBlockiä  aà  
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.
rD   rE   rG   rF   rI   .rH   r€   rå   rç   ræ   râ   Úinput_grid_sizeÚn_layersÚp_stochasticr   Nc                 óp  >• [         TU ]  5         [        U5      U:X  d  [        SU SU S35      e[        R
                  " 5       U l        [        USSSS9U l        [        U5       HL  u  nnUS:X  a  SOSnU =R                  [        US:X  a  UOUUUUUUUUUU	U
UU R                  US	9/-  sl        MN     g )
Nz'p_stochastic must have length n_layers=z, got p_stochastic=Ú.r(   r    r!   rM   r   )rD   rE   rG   rF   r   rI   rH   r€   rå   rç   ræ   râ   rä   rJ   )rY   rZ   Úlenrˆ   r	   Ú
ModuleListrc   r$   rä   Ú	enumerater  )rd   rD   rE   rG   rF   rI   rH   r€   rå   rç   ræ   râ   r  r  r  ÚidxrË   r   ri   s                     €r#   rZ   Ú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%   rk   c                 ó<   • U R                    H  nU" U5      nM     U$ r	  r
  )rd   rk   Úlayers      r#   ro   ÚMaxVitBlock.forward-  s    € ð —[”[ˆEÙ�a“ŠAñ !àˆr%   )rä   rc   )rr   rs   rt   ru   rv   r`   rw   r   r	   rx   r   ÚlistrZ   r
   ro   ry   rz   r{   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            !       ó  ^ • \ rS rSrSrS\R                  SSSSSS4S\\\4   S	\S
\S\	\   S\	\   S\S\
S\\S\R                  4      S\S\R                  4   S\
S\
S\S\
S\
S\SS4 U 4S jjjrS\S\4S jrS rSrU =r$ )r   i9  a1  
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£   rq   iè  r   Ústem_channelsrâ   Úblock_channelsÚblock_layersr€   Ústochastic_depth_probrI   .rH   rG   rF   rå   rç   ræ   Únum_classesr   c                 ó   >• [         TU ]  5         [        U 5        SnUc  [        [        R
                  SSS9n[        U[        U5      5      n[        U5       H6  u  nnUS   U-  S:w  d  US   U-  S:w  d  M   [        SU SU S	U S
U S3	5      e   [        R                  " [        UUSSUU	SS S9[        X"SSS S SS95      U l        [        USSSS9nX0l        [        R                  " 5       U l        U/US S -   nUn["        R$                  " SU['        U5      5      R)                  5       nSn[+        UUU5       HZ  u  nnnU R                   R-                  [/        UUU
UUU	UUUUUUUUUUU-    S95        U R                   S   R0                  nUU-  nM\     [        R                  " [        R2                  " S5      [        R4                  " 5       [        R6                  " US   5      [        R8                  " US   US   5      [        R:                  " 5       [        R8                  " US   USS95      U l        U R?                  5         g )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   rI   rH   rL   rR   T)r   rI   rH   rL   rM   r4   )rD   rE   rG   rF   rI   rH   r€   rå   rç   ræ   râ   r  r  r  )rL   ) rY   rZ   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)rd   r   r  râ   r  r  r€   r   rI   rH   rG   rF   rå   rç   ræ   r!  Úinput_channelsÚblock_input_sizesr  Úblock_input_sizerD   rE   r  Úp_idxÚ
in_channelÚout_channelÚ
num_layersri   s                              €r#   rZ   Ú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óð ñ &Bô —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ñ) 4`ô0 Ÿ-š-Ü× Ò  Ó#Ü�JŠJ‹LÜ�LŠL˜¨Ñ+Ó,Ü�IŠI�n RÑ(¨.¸Ñ*<Ó=Ü�GŠG‹IÜ�IŠI�n RÑ(¨+¸EÑBó
ˆŒð 	×ÑÕr%   rk   c                 ó€   • U R                  U5      nU R                   H  nU" U5      nM     U R                  U5      nU$ r¾   )r&  r'  r/  )rd   rk   Úblocks      r#   ro   ÚMaxVit.forwardÄ  s9   € Ø�I‰I�a‹LˆØ—[”[ˆEÙ�a“ŠAñ !à�O‰O˜AÓˆØˆr%   c                 ó(  • U R                  5        GH}  n[        U[        R                  5      (       ab  [        R                  R                  UR                  SS9  UR                  b+  [        R                  R                  UR                  5        Mƒ  M…  [        U[        R                  5      (       aV  [        R                  R                  UR                  S5        [        R                  R                  UR                  S5        Mú  [        U[        R                  5      (       d  GM  [        R                  R                  UR                  SS9  UR                  c  GMT  [        R                  R                  UR                  5        GM€     g )Nr…   r†   r!   r   )ÚmodulesÚ
isinstancer	   r[   r—   Únormal_ÚweightrL   Úzeros_r%  Ú	constant_r�   )rd   Úms     r#   r0  Ú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%   )r'  r/  râ   r&  )rr   rs   rt   ru   rv   r	   ÚGELUr   r`   r  rw   r   r   rx   rZ   r
   ro   r0  ry   rz   r{   s   @r#   r   r   9  s0  ø† ñð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ð tðl˜ð  Fô ÷+ð +r%   r   r  r  r  r   râ   r€   ÚweightsÚprogressÚkwargsc                 ód  • Ube  [        US[        UR                  S   5      5        UR                  S   S   UR                  S   S   :X  d   e[        USUR                  S   5        UR                  SS5      n	[	        SU UUUUUU	S.UD6n
Ub  U
R                  UR                  US	S
95        U
$ )Nr!  Ú
categoriesÚmin_sizer   r!   r   ©éà   rM  )r  r  r  r   r€   râ   r   T)rG  Ú
check_hashr"   )r   r  ÚmetaÚpopr   Úload_state_dictÚget_state_dict)r  r  r  r   râ   r€   rF  rG  rH  r   Úmodels              r#   Ú_maxvitrT  Ú  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                   óf   • \ rS rSr\" S\" \SS\R                  S9\	SSSSS	S
S.0SSSS.S9r
\
rSrg)r   i  z9https://download.pytorch.org/models/maxvit_t-bc5ab103.pthrM  )Ú	crop_sizeÚresize_sizeÚinterpolationiÈË×rL  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.)rJ  Ú
num_paramsrK  ÚrecipeÚ_metricsÚ_opsÚ
_file_sizeÚ_docs)ÚurlÚ
transformsrO  r"   N)rr   rs   rt   ru   r   r   r   r   ÚBICUBICr   ÚIMAGENET1K_V1ÚDEFAULTry   r"   r%   r#   r   r     sb   † ÙàGÙØ¨3¸CÐO`×OhÑOhñ
ð /Ø"Ø"ØdàØ#Ø#ñ ðð Ø!ðgñ
ñ€Mð. ƒGr%   r   Ú
pretrained)rF  T)rF  rG  c                 ó\   • [         R                  U 5      n [        SS/ SQ/ SQSSSU US.UD6$ )	aF  
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:
é@   )rf  é€   é   i   )r    r    rÍ   r    é    gš™™™™™É?é   )r  r  r  r€   r   râ   rF  rG  r"   )r   ÚverifyrT  )rF  rG  rH  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@   rx   rB   r}   r¹   rÄ   rÕ   rá   r  r  r   rw   ÚboolrT  r   rb  r   r"   r%   r#   Ú<module>r|     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%   