
    ^j                    ^   d Z ddlZddl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mZmZmZmZ ddlZddlmZ ddlmZ dd	lmZmZ dd
lmZmZmZmZmZm Z m!Z!m"Z"m#Z#m$Z$m%Z%m&Z&m'Z'm(Z(m)Z)m*Z*m+Z+m,Z,m-Z-m.Z.m/Z/m0Z0m1Z1m2Z2m3Z3 ddl4m5Z5 ddl6m7Z7 ddl8m9Z9 ddl:m;Z;m<Z< ddl=m>Z>m?Z? g dZ@e G d d             ZAe G d d             ZBe G d d             ZC G d dej                        ZE G d dej                        ZF G d dej                        ZGddej                  d eHd!eHd"dfd#ZI G d$ d%ej                        ZJddej                  d eHd!eHd"dfd&ZKd'eeL   d(eLd"eLfd)ZM G d* d+ej                        ZN G d, d-ej                        ZOd.ej                  d/eeL   d"ej                  fd0ZQe9d1ej                  d/eeL   d2eeL   d"ej                  fd3       ZRd.ej                  d4eeL   d"ej                  fd5ZSe9d1ej                  d4eeL   d2eeL   d"ej                  fd6       ZTd7eAd/eeLeLf   d"ee   fd8ZU G d9 d:ej                        ZV G d; d<ej                        ZWd.ej                  d/eeL   d"ej                  fd=ZXe9d1ej                  d/eeL   d2eeL   d"ej                  fd>       ZYd.ej                  d4eeL   d"ej                  fd?ZZe9d1ej                  d4eeL   d2eeL   d"ej                  fd@       Z[ G dA dBej                        Z\ G dC dDej                        Z] G dE dFej                        Z^ G dG dHej                        Z_ G dI dJej                        Z`d7eAd2eeLeLf   d"eAfdKZad7eCdLed"eCfdMZb G dN dOej                        Zc	 	 	 	 	 	 	 	 	 	 	 	 ddUeHdVeHdWeddXeddYeHdZeHd[edd\eHd]eHd^eee   d_eHd`eLd"eeHef   fdaZf	 	 	 	 	 	 	 	 	 	 	 	 ddUeHdVeHdWeddceedZeHd\eHd]eHd/eeeLeLf      ddeLd^eee   d_eHd`eLd"eeHef   fdeZg	 	 	 	 	 	 	 	 	 	 	 ddUeHdVeHdZeHdheHd\eHd]eHd/eeeLeLf      diedd^eeeeeeeef   f   d_eHd`eLd"eeHef   fdjZhd"eeHef   fdkZi ejdi dl eCddmdndodp egdRdqr      ds eCddmdtdodp egdPdRdqu      dv eCddwdxdodp efdRdQy      dz eCddwd{dodp efdbdRdQ|      d} eCdd~d{ddp efdbd      d eCddd{ddp efdbddf      d eCddwdxdodp efdbdRdQd      d eCddmdtdodp egdRdqdgd      d eCddwdxdodp efdbdg      d eCddwd{dodp efddRdQdgd      d eCddwd{dodp efdbdgdT      d eCdd~d{ddp efdbddfdg      d eCddd{ddp efdbddfdg      d eCddmdtdodd ef       d eCddmdtdodd ehdSd      d eCdwdndd      d eCdwd{dd      d eCd~d{dd      d eCdd{dd      d eCdddd      d eCdddd      d eCdddddd eg       d eCddmdddod eg       d eCddmdddod eg       d eCddmdddod eg       d eCdddddd egdg      d eCddmdddod egdg      d eCddmdddod egdg      d eCddwdddod egdgdf      d eCddwd{ddodd egdg      d eCddmdddodd eh       d eCddmdddod eh       d eCddwdddd eh       d eCddwddddd ehdRdSƫ      d eCdd~dddd ehdRɫ      d eCddddddd ehdRɫ      d eCddmddddRdTdМ ei       d eCddwddddRddМ ei       d eCddwd{dddRddМ ei       d eCdd~d{dddRddМ ei       d eCddd{dddRddМ ei       ZkdeeHej                  f   dej                  d"eeHej                  f   fdׄZlddeHdeeH   deddLed"ecf
dۄZmddeHdLed"eeHef   fd݄Zn e>i d end߫      d enddd      d endd      d endd      d end      d end      d end      d enddddd      d enddeed      d enddd      d end߫      d endd      d endd      d end߫      d  end߫      d enddd      d endd      i d endd      d endd      d endd      d	 end߫      d
 end߫      d end߫      d end߫      d end߫      d end߫      d enddd      d endddd      d endd      d enddd      d enddd      d endddd      d endddd      d endddd      i d  endd!d      d" enddd      d# end      d$ enddddd      d% endd      d& endd'dd      d( enddd      d) endd*dd      d+ enddd,      d- end      d. enddddd      d/ end߫      d0 endd      d1 endee2      d3 enddddd      d4 endd5d6dd      d7 endee2      i d8 enddddd      d9 endd5d6dd      d: endee2      d; enddddd      d< endd5d6dd      d= endee2      d> enddddd      d? endd5d6dd      d@ enddA      dB enddddd      dC endd5d6dd      dD enddA      dE enddddd      dF endd5ddG      dH enddA      dI enddddd      dJ endd5d6dd            Zoe?ddeddLed"ecfdK       Zpe?ddeddLed"ecfdL       Zqe?ddeddLed"ecfdM       Zre?ddeddLed"ecfdN       Zse?ddeddLed"ecfdO       Zte?ddeddLed"ecfdP       Zue?ddeddLed"ecfdQ       Zve?ddeddLed"ecfdR       Zwe?ddeddLed"ecfdS       Zxe?ddeddLed"ecfdT       Zye?ddeddLed"ecfdU       Zze?ddeddLed"ecfdV       Z{e?ddeddLed"ecfdW       Z|e?ddeddLed"ecfdX       Z}e?ddeddLed"ecfdY       Z~e?ddeddLed"ecfdZ       Ze?ddeddLed"ecfd[       Ze?ddeddLed"ecfd\       Ze?ddeddLed"ecfd]       Ze?ddeddLed"ecfd^       Ze?ddeddLed"ecfd_       Ze?ddeddLed"ecfd`       Ze?ddeddLed"ecfda       Ze?ddeddLed"ecfdb       Ze?ddeddLed"ecfdc       Ze?ddeddLed"ecfdd       Ze?ddeddLed"ecfde       Ze?ddeddLed"ecfdf       Ze?ddeddLed"ecfdg       Ze?ddeddLed"ecfdh       Ze?ddeddLed"ecfdi       Ze?ddeddLed"ecfdj       Ze?ddeddLed"ecfdk       Ze?ddeddLed"ecfdl       Ze?ddeddLed"ecfdm       Ze?ddeddLed"ecfdn       Ze?ddeddLed"ecfdo       Ze?ddeddLed"ecfdp       Ze?ddeddLed"ecfdq       Ze?ddeddLed"ecfdr       Ze?ddeddLed"ecfds       Ze?ddeddLed"ecfdt       Ze?ddeddLed"ecfdu       Ze?ddeddLed"ecfdv       Ze?ddeddLed"ecfdw       Ze?ddeddLed"ecfdx       Ze?ddeddLed"ecfdy       Ze?ddeddLed"ecfdz       Ze?ddeddLed"ecfd{       Ze?ddeddLed"ecfd|       Ze?ddeddLed"ecfd}       Ze?ddeddLed"ecfd~       Ze?ddeddLed"ecfd       Ze?ddeddLed"ecfd       Ze?ddeddLed"ecfd       Ze?ddeddLed"ecfd       Zy(  a   MaxVit and CoAtNet Vision Transformer - CNN Hybrids in PyTorch

This is a from-scratch implementation of both CoAtNet and MaxVit in PyTorch.

99% of the implementation was done from papers, however last minute some adjustments were made
based on the (as yet unfinished?) public code release https://github.com/google-research/maxvit

There are multiple sets of models defined for both architectures. Typically, names with a
 `_rw` suffix are my own original configs prior to referencing https://github.com/google-research/maxvit.
These configs work well and appear to be a bit faster / lower resource than the paper.

The models without extra prefix / suffix' (coatnet_0_224, maxvit_tiny_224, etc), are intended to
match paper, BUT, without any official pretrained weights it's difficult to confirm a 100% match.

Papers:

MaxViT: Multi-Axis Vision Transformer - https://arxiv.org/abs/2204.01697
@article{tu2022maxvit,
  title={MaxViT: Multi-Axis Vision Transformer},
  author={Tu, Zhengzhong and Talebi, Hossein and Zhang, Han and Yang, Feng and Milanfar, Peyman and Bovik, Alan and Li, Yinxiao},
  journal={ECCV},
  year={2022},
}

CoAtNet: Marrying Convolution and Attention for All Data Sizes - https://arxiv.org/abs/2106.04803
@article{DBLP:journals/corr/abs-2106-04803,
  author    = {Zihang Dai and Hanxiao Liu and Quoc V. Le and Mingxing Tan},
  title     = {CoAtNet: Marrying Convolution and Attention for All Data Sizes},
  journal   = {CoRR},
  volume    = {abs/2106.04803},
  year      = {2021}
}

Hacked together by / Copyright 2022, Ross Wightman
    N)OrderedDict)	dataclassreplacefield)partial)AnyCallableDictListOptionalSetTupleUnion)nn)Final)IMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)MlpConvMlpDropPathcalculate_drop_path_rates	LayerNorm
LayerScaleLayerScale2dClassifierHeadNormMlpClassifierHeadcreate_attnget_act_layerget_norm_layerget_norm_act_layercreate_conv2dcreate_pool2dtrunc_normal_tf_	to_2tupleextend_tuplemake_divisible_assert	RelPosMlp
RelPosBiasRelPosBiasTfuse_fused_attnresize_rel_pos_bias_table   )build_model_with_cfg)feature_take_indices)register_notrace_function)named_applycheckpoint_seq)generate_default_cfgsregister_model)
MaxxVitCfgMaxxVitConvCfgMaxxVitTransformerCfgMaxxVitc                   d   e Zd ZU dZdZeed<   dZeed<   dZ	e
ed<   dZeed<   dZeed	<   dZeed
<   dZe
ed<   dZe
ed<   dZeed<   dZeed<   dZeed<   dZeed<   dZeeeef      ed<   dZeeeef      ed<   dZeed<   dZeed<   dZee
   ed<   dZeed<   dZeed<   d Zeed!<   d"Ze
ed#<   d$ Z y)%r7   z-Configuration for MaxxVit transformer blocks.    dim_headT
head_first      @expand_ratioexpand_firstshortcut_bias	attn_bias        	attn_drop	proj_dropavg2	pool_typebiasrel_pos_type   rel_pos_dimpartition_ratioNwindow_size	grid_sizeFno_block_attnuse_nchw_attninit_valuesgelu	act_layerlayernorm2d
norm_layer	layernormnorm_layer_clư>norm_epsc                     | j                   t        | j                         | _         | j                  9t        | j                        | _        | j                   | j                  | _         y y y N)rM   r$   rL   selfs    ^/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/maxxvit.py__post_init__z#MaxxVitTransformerCfg.__post_init__m   s\    >>%&t~~6DN'()9)9:D~~%!%!1!1 & (    )!__name__
__module____qualname____doc__r;   int__annotations__r<   boolr>   floatr?   r@   rA   rC   rD   rF   strrH   rJ   rK   rL   r   r   rM   rN   rO   rP   rR   rT   rV   rX   r^    r_   r]   r7   r7   T   s    7HcJL%L$M4ItIuIuIsL#KOS-1K%S/*1+/Ixc3h(/M4M4#'K%'Is#J#$M3$He2r_   r7   c                   <   e Zd ZU dZdZeed<   dZeed<   dZ	e
ed<   dZeed	<   d
Zeed<   dZe
ed<   dZe
ed<   dZeed<   dZeed<   dZeed<   dZeed<   dZe
ed<   dZeed<   dZeed<   dZeed<   dZee   ed<   dZeed <   dZeed!<   dZeed"<   d#Zee   ed$<   d% Zy#)&r6   z-Configuration for MaxxVit convolution blocks.mbconv
block_typer=   r>   Texpand_output   kernel_sizer-   
group_sizeFpre_norm_actoutput_biasdwstride_moderE   rF   downsample_pool_type padding
attn_earlyse
attn_layersiluattn_act_layer      ?
attn_ratiorW   rP   rQ   rR   rT   rV   NrX   c                    | j                   dv sJ | j                   dk(  }| j                  s|rdnd| _        | j                  s	|sd| _        | j                  |rdnd| _        | j                  xs | j
                  | _        y )N)rk   convnextrk   batchnorm2drS   rU   h㈵>rW   )rl   rT   rV   rX   ru   rF   )r\   
use_mbconvs     r]   r^   zMaxxVitConvCfg.__post_init__   st    "8888__0
/9m}DO!!*!,D== $.DDDM$($=$=$O!r_   )r`   ra   rb   rc   rl   rh   re   r>   rg   rm   rf   ro   rd   rp   rq   rr   rt   rF   ru   rw   rx   rz   r|   r~   rP   r   rR   rT   rV   rX   r^   ri   r_   r]   r6   r6   v   s    7JL%M4KJL$KKIs &#&GSJJ NC J#'K%'IsJM3 $Hhuo$
Pr_   r6   c                       e Zd ZU dZdZeedf   ed<   dZeedf   ed<   dZ	ee
eeedf   f   df   ed<   d	Ze
eeeef   f   ed
<   dZeed<    ee      Zeed<    ee      Zeed<   dZee   ed<   dZeed<   y)r5   z!Configuration for MaxxVit models.`           .	embed_dim   rn      r   depths)Cr   Tr   rl   @   
stem_widthF	stem_bias)default_factoryconv_cfgtransformer_cfgNhead_hidden_sizevit_effweight_init)r`   ra   rb   rc   r   r   rd   re   r   rl   r   rh   r   r   rf   r   r6   r   r7   r   r   r   r   ri   r_   r]   r5   r5      s    +!4IuS#X4*FE#s(O*:NJeCsCx01367N.0Jc5c?*+0It$^DHnD-2CX-YO*Y&*hsm* K r_   r5   c                        e Zd ZU dZee   ed<   	 	 	 	 	 	 	 	 	 	 ddedee   dedededed	ee	   d
e
de
f fdZddej                  deej                     dej                  fdZ xZS )Attention2dz)Multi-head attention for 2D NCHW tensors.
fused_attndimdim_outr;   rG   r?   r<   rel_pos_clsrC   rD   c                    |
|d}t         |           |xs |}|r|n|}||z  | _        || _        || _        |dz  | _        t               | _        t        j                  ||dz  dfd|i|| _
        |r |dd| j                  i|nd| _        t        j                  |      | _        t        j                  ||dfd|i|| _        t        j                  |	      | _        y)	  
        Args:
            dim: Input dimension.
            dim_out: Output dimension (defaults to input dimension).
            dim_head: Dimension per attention head.
            bias: Whether to use bias in qkv and projection.
            expand_first: Whether to expand channels before or after qkv.
            head_first: Whether heads are first in tensor layout.
            rel_pos_cls: Relative position class to use.
            attn_drop: Attention dropout rate.
            proj_drop: Projection dropout rate.
        devicedtype      rn   r-   rG   	num_headsNri   )super__init__r   r;   r<   scaler+   r   r   Conv2dqkvrel_posDropoutrC   projrD   r\   r   r   r;   rG   r?   r<   r   rC   rD   r   r   dddim_attn	__class__s                 r]   r   zAttention2d.__init__   s    4 /.S*7!X- $%
(*99S(Q,CCCFQ{BT^^BrBW[I.IIhDDD	I.r_   xshared_rel_posreturnc                    |j                   \  }}}}| j                  rP| j                  |      j                  || j                  | j
                  dz  d      j                  dd      \  }}}	nK| j                  |      j                  |d| j                  | j
                  d      j                  d      \  }}}	| j                  rd }
| j                  | j                  j                         }
n||}
t        j                  j                  j                  |j!                  dd      j#                         |j!                  dd      j#                         |	j!                  dd      j#                         |
| j$                  r| j&                  j(                  nd      j!                  dd      j                  |d||      }n|| j*                  z  }|j!                  dd      |z  }| j                  | j                  |      }n|||z   }|j-                  d      }| j'                  |      }|	|j!                  dd      z  j                  |d||      }| j/                  |      }| j1                  |      }|S )	Nrn   r   r   r-   rB   	attn_mask	dropout_p)shaper<   r   viewr   r;   chunkreshapeunbindr   r   get_biastorchr   
functionalscaled_dot_product_attention	transpose
contiguoustrainingrC   pr   softmaxr   rD   )r\   r   r   Br   HWqkvrA   attns               r]   forwardzAttention2d.forward   s   WW
1a??hhqk&&q$..$--!:KRPVVWX^_V`GAq!hhqk))!QrRYYZ[\GAq!??I||' LL113	+*	##@@B#..0B#..0B#..0#.2mm$..** A  iB2q! 4  DJJA;;r2&*D||'||D)+n,<<B<'D>>$'DT^^B++11!RA>AIIaLNN1r_   
Nr:   TTTNrB   rB   NNrZ   r`   ra   rb   rc   r   rf   re   rd   r   r	   rg   r   r   Tensorr   __classcell__r   s   @r]   r   r      s    3d
 &*!%#.2!!(/(/ c](/ 	(/
 (/ (/ (/ "(+(/ (/ (/T# #x7M #Y^YeYe #r_   r   c                        e Zd ZU dZee   ed<   	 	 	 	 	 	 	 	 	 	 ddedee   dedededed	ee	   d
e
de
f fdZddej                  deej                     dej                  fdZ xZS )AttentionClz/Channels-last multi-head attention (B, ..., C).r   r   r   r;   rG   r?   r<   r   rC   rD   c                    |
|d}t         |           |xs |}|r||kD  r|n|}||z  dk(  sJ d       ||z  | _        || _        || _        |dz  | _        t               | _        t        j                  ||dz  fd|i|| _
        |r |d	d| j                  i|nd| _        t        j                  |      | _        t        j                  ||fd|i|| _        t        j                  |	      | _        y)
r   r   r   z(attn dim should be divisible by head_dimr   rn   rG   r   Nri   )r   r   r   r;   r<   r   r+   r   r   Linearr   r   r   rC   r   rD   r   s                 r]   r   zAttentionCl.__init__  s    4 /.S*w}7#("a'S)SS'!X- $%
(*99S(Q,@T@R@FQ{BT^^BrBW[I.IIhAdAbA	I.r_   r   r   r   c                 V   |j                   d   }|j                   d d }| j                  r`| j                  |      j                  |d| j                  | j
                  dz        j                  dd      j                  dd      \  }}}n[| j                  |      j                  |dd| j                  | j
                        j                  dd      j                  d      \  }}}| j                  r~d }| j                  | j                  j                         }n||}t        j                  j                  j!                  ||||| j"                  r| j$                  j&                  nd      }ns|| j(                  z  }||j                  d	d      z  }	| j                  | j                  |	|
      }	n||	|z   }	|	j+                  d      }	| j%                  |	      }	|	|z  }|j                  dd      j                  |dz         }| j-                  |      }| j/                  |      }|S )Nr   r   rn   r-   r   r   rB   r   r   r   )r   )r   r<   r   r   r   r;   r   r   r   r   r   r   r   r   r   r   r   r   rC   r   r   r   r   rD   )
r\   r   r   r   restore_shaper   r   r   rA   r   s
             r]   r   zAttentionCl.forward,  s   GGAJ??hhqk&&q"dnndmma>OPZZ[\^_`ffghnofpGAq!hhqk))!RDNNDMMR\\]^`abiijklGAq!??I||' LL113	+*	##@@1a#.2mm$..** A A DJJAq{{2r**D||'||D|H+n,<<B<'D>>$'DqAKK1%%me&;<IIaLNN1r_   r   rZ   r   r   s   @r]   r   r      s    9d
 &*!%#.2!!)/)/ c])/ 	)/
 )/ )/ )/ "(+)/ )/ )/V# #x7M #Y^YeYe #r_   r   c                   |     e Zd ZdZ	 	 	 	 	 ddededededef
 fdZdej                  d	ej                  fd
Z
 xZS )Downsample2da5  A downsample pooling module supporting several maxpool and avgpool modes.

    * 'max' - MaxPool2d w/ kernel_size 3, stride 2, padding 1
    * 'max2' - MaxPool2d w/ kernel_size = stride = 2
    * 'avg' - AvgPool2d w/ kernel_size 3, stride 2, padding 1
    * 'avg2' - AvgPool2d w/ kernel_size = stride = 2
    r   r   rF   rw   rG   c                    t         |           |dv sJ |dk(  rt        ddd|xs d      | _        nS|dk(  rt        dd|xs d	      | _        n6|d
k(  rt        d
ddd|xs d      | _        nt        d
d|xs d	      | _        ||k7  r!t	        j
                  ||d|||      | _        yt	        j                         | _        y)z
        Args:
            dim: Input dimension.
            dim_out: Output dimension.
            pool_type: Type of pooling operation.
            padding: Padding mode.
            bias: Whether to use bias in expansion conv.
        )maxmax2avgrE   r   rn   r   r-   )ro   striderw   r   r   )rw   r   F)ro   r   count_include_padrw   )rG   r   r   N)r   r   r"   poolr   r   expandIdentity)	r\   r   r   rF   rw   rG   r   r   r   s	           r]   r   zDownsample2d.__init__[  s    $ 	::::%e1glYZ[DI& %eQ1EDI%%1Q%QXQ]\]_DI &eQ1EDI'>))C!$vUZ[DK++-DKr_   r   r   c                 J    | j                  |      }| j                  |      }|S rZ   )r   r   r\   r   s     r]   r   zDownsample2d.forward~  s!    IIaLKKNr_   )rE   rv   TNN)r`   ra   rb   rc   rd   rh   rf   r   r   r   r   r   r   s   @r]   r   r   R  sj     $!(!( !( 	!(
 !( !(F %,, r_   r   rv   modulenameschemer   c                    t        | t        j                  t        j                  f      r|dk(  rbt        j                  j                  | j                  d       | j                  *t        j                  j                  | j                         yy|dk(  rNt        | j                  d       | j                  *t        j                  j                  | j                         yy|dk(  r`t        j                  j                  | j                         | j                  *t        j                  j                  | j                         yyt        j                  j                  | j                         | j                  Zd|v r,t        j                  j                  | j                  d       yt        j                  j                  | j                         yyy)	z&Initialize transformer module weights.normal{Gz?stdNtrunc_normalxavier_normalmlprW   )
isinstancer   r   r   initnormal_weightrG   zeros_r#   xavier_normal_xavier_uniform_)r   r   r   s      r]   _init_transformerr     s5   &299bii01XGGOOFMMsO3{{&v{{+ '~%V]]4{{&v{{+ '&GG""6==1{{&v{{+ ' GG##FMM2{{&D=GGOOFKKTO:GGNN6;;/	 '! 2r_   c                        e Zd ZdZdd e       dddfdedededee   d	ed
ef fdZ	dde
ddfdZddej                  deej                     dej                  fdZ xZS )TransformerBlock2daY  Transformer block with 2D downsampling.

    '2D' NCHW tensor layout

    Some gains can be seen on GPU using a 1D / CL block, BUT w/ the need to switch back/forth to NCHW
    for spatial pooling, the benefit is minimal so ended up using just this variant for CoAt configs.

    This impl was faster on TPU w/ PT XLA than the 1D experiment.
    r-   NrB   r   r   r   r   cfg	drop_pathc	                    ||d}	t         |           t        t        |j                        |j
                        }
t        |j                        }|dk(  rqt        ||f|j                  |j                  d|	| _        t        j                  t        d |
|fi |	fdt        ||fd|j                  i|	fg            | _        n.||k(  sJ t        j                          | _         |
|fi |	| _        t#        ||f|j$                  |j&                  |j(                  ||j*                  |j,                  d|	| _        |j0                  rt3        |fd	|j0                  i|	nt        j                          | _        |d
kD  rt7        |      nt        j                          | _         |
|fi |	| _        t=        d|t?        ||j@                  z        ||j,                  d|	| _!        |j0                  rt3        |fd	|j0                  i|	nt        j                          | _"        |d
kD  rt7        |      | _#        yt        j                          | _#        y)a  
        Args:
            dim: Input dimension.
            dim_out: Output dimension.
            stride: Stride for downsampling.
            rel_pos_cls: Relative position class.
            cfg: Transformer block configuration.
            drop_path: Drop path rate.
        r   epsr   )rF   rG   normdownrF   )r;   r?   rG   r   rC   rD   rP   rB   in_featureshidden_featuresrR   dropNri   )$r   r   r   r   rT   rX   r   rR   r   rF   r@   shortcutr   
Sequentialr   norm1r   r   r;   r?   rA   rC   rD   r   rP   r   ls1r   
drop_path1norm2r   rd   r>   r   ls2
drop_path2)r\   r   r   r   r   r  r  r   r   r   rT   rR   r   s               r]   r   zTransformerBlock2d.__init__  s   ( /^CNN;N
!#--0	Q;(gmUXUfUfmjlmDM{C.2./c3N#--N2NO4 ( DJ
 '>!>KKMDM#C.2.DJ

 \\))#mmmm

 

	 PS<KS__KKdfdododq1:R(9-R[[].2.
 
#*:*: :;	

 
 PS<KS__KKdfdododq1:R(9-R[[]r_   r   r   c                 :    t        t        t        |      |        y Nr   )r1   r   r   r\   r   s     r]   init_weightszTransformerBlock2d.init_weights  s    G-f=tDr_   r   r   c           
      ,   | j                  |      | j                  | j                  | j                  | j	                  |      |                  z   }|| j                  | j                  | j                  | j                  |                        z   }|S )Nr   )	r  r  r  r   r  r  r  r   r  )r\   r   r   s      r]   r   zTransformerBlock2d.forward  so    MM!ttxx		$**Q-`n	8o/pqq$**Q-)@ ABBr_   rv   rZ   )r`   ra   rb   rc   r7   rd   r   r	   rg   r   rh   r  r   r   r   r   r   s   @r]   r   r     s     .2)>)@!;S;S ;S 	;S
 "(+;S ';S ;SzE3 E E x7M Y^YeYe r_   r   c                    t        | t        j                        r|dk(  rbt        j                  j	                  | j
                  d       | j                  *t        j                  j                  | j                         yy|dk(  rNt        | j
                  d       | j                  *t        j                  j                  | j                         yy|dk(  r`t        j                  j                  | j
                         | j                  *t        j                  j                  | j                         yy| j                  d   | j                  d   z  | j                  z  }|| j                  z  }t        j                  j	                  | j
                  dt        j                  d	|z               | j                  *t        j                  j                  | j                         yyy)
z&Initialize convolution module weights.r   r   r   Nr   r   r   r-   g       @)r   r   r   r   r   r   rG   r   r#   r   ro   out_channelsgroupsmathsqrt)r   r   r   fan_outs       r]   
_init_convr!    sU   &"))$XGGOOFMMsO3{{&v{{+ '~%V]]4{{&v{{+ '&GG""6==1{{&v{{+ ' ((+f.@.@.CCfFYFYYG%GGGOOFMM1diig.FG{{&v{{+ '% %r_   rp   channelsc                 &    | sy|| z  dk(  sJ || z  S )z3Calculate number of groups for grouped convolution.r-   r   ri   )rp   r"  s     r]   
num_groupsr$    s(     *$))):%%r_   c                        e Zd ZdZdd e       dddfdededed	eeef   d
edef fdZdde	ddfdZ
dej                  dej                  fdZ xZS )MbConvBlockzGPre-Norm Conv Block - 1x1 - kxk - 1x1, w/ inverted bottleneck (expand).r-   r-   r-   rB   Nin_chsout_chsr   dilationr  r  c	                    ||d}	t         |           t        t        |j                  |j
                        |j                        }
t        |j                  r|n||j                  z        }t        |j                  |      }|dk(  r5t        ||f|j                  |j                  |j                  d|	| _        nt#        j$                         | _        |j&                  dv sJ d\  }}}|j&                  dk(  r||d   }}n|j&                  d	k(  r||d   }}n||d
   }} |
|fd|j(                  i|	| _        |dkD  r*t        ||f|j,                  |j                  d|	| _        nt#        j$                         | _        t1        ||dfd|i|	| _         |
|fi |	| _        t1        |||j6                  f||||j                  d|	| _        i }t;        |j<                  t>              rV|j<                  dk(  s|j<                  dk(  r8|j@                  |d<   tC        |jD                  |j                  r|n|z        |d<   |jF                  r4tI        |j<                  |fi ||	| _%         |
|fi |	| _&        d| _'        n3d| _%         |
|fi |	| _&        tI        |j<                  |fi ||	| _'        t1        ||dfd|j                  i|	| _(        |dkD  rtS        |      | _*        yt#        j$                         | _*        y)a  
        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            stride: Stride for conv.
            dilation: Dilation for conv.
            cfg: Convolution block configuration.
            drop_path: Drop path rate.
        r   r  r   )rF   rG   rw   )r   1x1rs   )r-   r-   r-   r   r-   r,  r   	apply_act)rF   rw   r   )r   r*  r  rw   ry   ecarR   rd_channelsNrG   rB   )+r   r   r   r    rT   rR   rX   r&   rm   r>   r$  rp   r   rF   rr   rw   r  r   r   rt   rq   pre_normru   r  r!   	conv1_1x1r  ro   	conv2_kxkr   rz   rh   r|   rd   r~   rx   r   se_earlyr  ry   	conv3_1x1r   r  )r\   r(  r)  r   r*  r  r  r   r   r   norm_act_layermid_chsr  stride_poolstride_1stride_2
dilation_2attn_kwargsr   s                     r]   r   zMbConvBlock.__init__  s   ( / !3CNNCMM!RX[XdXde S->->'FcN^N^!^_CNNG4Q;(k+.==sX[XcXckgikDM KKMDM"7777*1'Xx??f$&,hqkK__%#)8A;jH#)8A;jH&vP9I9IPRP?$VVss?W?WadalalsprsDIDI&vwQ(QbQ#G2r2
&OO	
 KK	
 	
 cnnc*~~%5)@+.+=+=K(-0cN_N_7el1m-nM* >>'U;URTUDM'626DJDG DM'626DJ!#..'O[OBODG&wWWTVW09B),BKKMr_   r   r   c                 :    t        t        t        |      |        y r  r1   r   r!  r  s     r]   r  zMbConvBlock.init_weightse      GJv6=r_   r   c                    | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }| j                  | j                  |      }| j                  |      }| j                  | j                  |      }| j                  |      }| j                  |      |z   }|S rZ   )r  r0  r  r1  r  r2  r3  r  ry   r4  r  r\   r   r  s      r]   r   zMbConvBlock.forwardh  s    ==#MM!IIaL NN1JJqM NN1==$a AJJqM77
A NN1NN1(r_   r  )r`   ra   rb   rc   r6   rd   r   rg   r   rh   r  r   r   r   r   r   s   @r]   r&  r&    s    Q (."0"2!PRPR PR 	PR
 CHoPR  PR PRd>3 > > %,, r_   r&  c                        e Zd ZdZdddd e       ddddf	ded	ee   d
ededeeef   dedede	f fdZ
dej                  dej                  fdZ xZS )ConvNeXtBlockzConvNeXt Block.N   r-   r'  TrB   r(  r)  ro   r   r*  r  conv_mlpr  c           	         |	|
d}t         |           |xs |}t        |j                        }|r1t	        t        |j                        |j                        }t        }nd|j                  v sJ t        }t        }|| _        |dk(  rt        ||fi || _        nH||k7  r*t        j                  ||fd|j                   d|| _        nt        j"                         | _        |j$                  dv sJ d\  }}|j$                  d	k(  r|}n|}|dk(  rt        ||fd
|j&                  i|| _        nt        j"                         | _        t+        ||f|||d   d|j                   d|| _         ||fi || _         ||t1        |j2                  |z        f|j                   |d|| _        |r=|j6                  rt9        ||j6                  fi |nt        j"                         | _        n<|j6                  rt=        ||j6                  fi |nt        j"                         | _        |dkD  rt?        |      | _         yt        j"                         | _         y)ay  
        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            kernel_size: Kernel size for depthwise conv.
            stride: Stride for conv.
            dilation: Dilation for conv.
            cfg: Convolution block configuration.
            conv_mlp: Whether to use convolutional MLP.
            drop_path: Drop path rate.
        r   r  rU   r   r-   )ro   rG   )r   rs   r'  r   rF   T)ro   r   r*  	depthwiserG   )rG   rR   rB   N)!r   r   r   rR   r   r   rT   rX   r   r   r   use_conv_mlpr   r  r   r   rr   r   rt   ru   r  r!   conv_dwr  rd   r>   r   rP   r   lsr   r   r  )r\   r(  r)  ro   r   r*  r  rD  r  r   r   r   rR   rT   	mlp_layerr7  	stride_dwr   s                    r]   r   zConvNeXtBlock.__init__  s.   0 /#V!#--0	 !?S\\RJI#..000"JI$Q;(?B?DMwIIfga13??a^`aDMKKMDM.000!%Y??f$ KI!$VV^s?W?W^[]^DIDI$	
 $a[	
 	
 w-"-	  7*+
 	

 
 FIool7COOBrB[][f[f[hDGDGOOj#//@R@Y[YdYdYfDG09B),BKKMr_   r   r   c                    | j                  |      }| j                  |      }| j                  |      }| j                  r4| j	                  |      }| j                  |      }| j                  |      }n[|j                  dddd      }| j	                  |      }| j                  |      }| j                  |      }|j                  dddd      }| j                  |      |z   }|S Nr   r   rn   r-   )	r  r  rH  rG  r  r   rI  permuter  r@  s      r]   r   zConvNeXtBlock.forward  s    ==#IIaLLLO		!AA
A		!Q1%A		!AA
A		!Q1%ANN1(r_   )r`   ra   rb   rc   r6   rd   r   r   rf   rg   r   r   r   r   r   r   s   @r]   rB  rB    s    
 &* (."0"2!!OROR c]OR 	OR
 OR CHoOR  OR OR ORb %,, r_   rB  r   rL   c                 l   | j                   \  }}}}t        ||d   z  dk(  d| d|d    d       t        ||d   z  dk(  d| d|d    d       | j                  |||d   z  |d   ||d   z  |d   |      } | j                  ddddd	d
      j	                         j                  d|d   |d   |      }|S )z'Partition into non-overlapping windows.r   height () must be divisible by window ()r-   width (rn   r      r   r   r   r'   r   rN  r   )r   rL   r   r   r   r   windowss          r]   window_partitionrW    s    JAq!QAA!#xs2QR]^_R`Qaab%cdAA!#wqc1PQ\]^Q_P``a%bc	q!{1~%{1~qKN7JKXYN\]^Aii1aAq)446;;BAP[\]P^`abGNr_   rV  img_sizec                     |\  }}| j                   d   }| j                  d||d   z  ||d   z  |d   |d   |      }|j                  dddddd      j                         j                  d|||      }|S )zReverse window partition.r   r   r-   rn   r   rT  r   r   r   rN  r   rV  rL   rX  r   r   r   r   s          r]   window_reverser\    s     DAqbARk!n,a;q>.A;q>S^_`SacdeA			!Q1a#..055b!QBAHr_   rM   c           	      h   | j                   \  }}}}t        ||d   z  dk(  d| d|d           t        ||d   z  dk(  d| d|d           | j                  ||d   ||d   z  |d   ||d   z  |      } | j                  dddddd	      j	                         j                  d
|d   |d   |      }|S )z6Partition into overlapping windows with grid striding.r   height  must be divisible by grid r-   width r   rT  rn   r   r   rU  )r   rM   r   r   r   r   rV  s          r]   grid_partitionra    s    JAq!QA	!!WQC/J9UV<.#YZA	!!VA3.I)TU,#XY	q)A,Yq\ 19Q<iPQlARTUVAii1aAq)446;;B	!iXYl\]^GNr_   c                     |\  }}| j                   d   }| j                  d||d   z  ||d   z  |d   |d   |      }|j                  dddddd      j                         j                  d|||      }|S )zReverse grid partition.r   r   r-   rn   rT  r   r   rZ  rV  rM   rX  r   r   r   r   s          r]   grid_reverserd    s     DAqbARil*A1,=y|YWX\[\]A			!Q1a#..055b!QBAHr_   r  c                     d}| j                   dk(  rt        t        || j                        }|S | j                   dk(  rt        t        |      }|S | j                   dk(  rt        t
        |      }|S )z,Get relative position class based on config.Nr   )rL   
hidden_dimrG   )rL   bias_tf)rH   r   r(   rJ   r)   r*   )r  rL   r   s      r]   get_rel_pos_clsrh    su    K
5 i[S__]
 	 
		V	#jkB  
		Y	&lDr_   c            	       V     e Zd ZdZd e       dddfdedededef fd	Zd
 Z	d Z
 xZS )PartitionAttentionClzRGrid or Block partition + Attn + FFN.

    NxC 'channels last' tensor layout.
    blockrB   Nr   partition_typer  r  c           
      *   ||d}t         |           t        t        |j                        |j
                        }t        |j                        }	|dk(  | _        t        | j                  r|j                  n|j                        | _        t        || j                        }
 ||fi || _        t        ||f|j                   |j"                  |j$                  |
|j&                  |j(                  d|| _        |j,                  rt/        |fd|j,                  i|nt1        j2                         | _        |dkD  rt7        |      nt1        j2                         | _         ||fi || _        t=        d|t?        ||j@                  z        |	|j(                  d|| _!        |j,                  rt/        |fd|j,                  i|nt1        j2                         | _"        |dkD  rt7        |      | _#        y t1        j2                         | _#        y )	Nr   r  rk  r;   rG   r<   r   rC   rD   rP   rB   r  ri   )$r   r   r   r   rV   rX   r   rR   partition_blockr$   rL   rM   partition_sizerh  r  r   r;   rA   r<   rC   rD   r   rP   r   r   r   r  r   r  r  r   rd   r>   r   r  r  r\   r   rl  r  r  r   r   r   rT   rR   r   r   s              r]   r   zPartitionAttentionCl.__init__   s    /^C,=,=>CLLQ
!#--0	-8'4;O;OUXUbUbc%c4+>+>?*r*


 \\~~#mmmm

 

	 JM:cEsE"E^`^i^i^k1:R(9-R[[]*r*
 
c&6&6 67	

 
 JM:cEsE"E^`^i^i^k1:R(9-R[[]r_   c                 0   |j                   dd }| j                  rt        || j                        }nt	        || j                        }| j                  |      }| j                  rt        || j                  |      }|S t        || j                  |      }|S )Nr-   rn   )r   ro  rW  rp  ra  r   r\  rd  r\   r   rX  partitioneds       r]   _partition_attnz$PartitionAttentionCl._partition_attnL  s    771Q<*1d.A.ABK(D,?,?@Kii,{D,?,?JA  [$*=*=xHAr_   c           
      
   || j                  | j                  | j                  | j                  |                        z   }|| j	                  | j                  | j                  | j                  |                        z   }|S rZ   r  r  ru  r  r  r  r   r  r   s     r]   r   zPartitionAttentionCl.forward[  c    )=)=djjm)L MNN$**Q-)@ ABBr_   )r`   ra   rb   rc   r7   rd   rh   rg   r   ru  r   r   r   s   @r]   rj  rj    sW     #*)>)@!*S*S  *S '	*S
 *SXr_   rj  c                        e Zd ZdZ e       dddfdededef fdZdej                  d	ej                  fd
Z
dej                  d	ej                  fdZ xZS )ParallelPartitionAttentionzQExperimental. Grid and Block partition + single FFN.

    NxC tensor layout.
    rB   Nr   r  r  c           
         ||d}t         
|           |dz  dk(  sJ t        t        |j                        |j
                        }t        |j                        }|j                  |j                  k(  sJ t        |j                        | _        t        || j                        }	 ||fi || _        t        ||dz  f|j                  |j                   |j"                  |	|j$                  |j&                  d|| _        t        ||dz  f|j                  |j                   |j"                  |	|j$                  |j&                  d|| _        |j,                  rt/        |fd|j,                  i|nt1        j2                         | _        |dkD  rt7        |      nt1        j2                         | _         ||fi || _        t=        d
|t?        ||j@                  z        |||j&                  d|| _!        |j,                  rt/        |fd|j,                  i|nt1        j2                         | _"        |dkD  rt7        |      | _#        y	t1        j2                         | _#        y	)z
        Args:
            dim: Input dimension.
            cfg: Transformer block configuration.
            drop_path: Drop path rate.
        r   r   r   r  rn  rP   rB   )r	  r
  out_featuresrR   r  Nri   )$r   r   r   r   rV   rX   r   rR   rL   rM   r$   rp  rh  r  r   r;   rA   r<   rC   rD   
attn_block	attn_gridrP   r   r   r   r  r   r  r  r   rd   r>   r   r  r  )r\   r   r  r  r   r   r   rT   rR   r   r   s             r]   r   z#ParallelPartitionAttention.__init__g  s     /Qw!||^C,=,=>CLLQ
!#--0	#--///'8%c4+>+>?*r*
%1H

 \\~~#mmmm

 

 %1H

 \\~~#mmmm

 

 JM:cEsE"E^`^i^i^k1:R(9-R[[]*r*
 
c&6&6 67
 
 JM:cEsE"E^`^i^i^k1:R(9-R[[]r_   r   r   c                 J   |j                   dd }t        || j                        }| j                  |      }t	        || j                  |      }t        || j                        }| j                  |      }t        || j                  |      }t        j                  ||gd      S )Nr-   rn   r   r   )
r   rW  rp  r}  r\  ra  r~  rd  r   cat)r\   r   rX  partitioned_blockx_windowpartitioned_gridx_grids          r]   ru  z*ParallelPartitionAttention._partition_attn  s    771Q<,Q0C0CD OO,=>!"3T5H5H(S)!T-@-@A>>*:;.0C0CXNyy(F+44r_   c           
      
   || j                  | j                  | j                  | j                  |                        z   }|| j	                  | j                  | j                  | j                  |                        z   }|S rZ   rw  r   s     r]   r   z"ParallelPartitionAttention.forward  rx  r_   )r`   ra   rb   rc   r7   rd   rg   r   r   r   ru  r   r   r   s   @r]   rz  rz  a  su     *?)@!<S<S '<S 	<S|5 5%,, 5 %,, r_   rz  c           	      l   | j                   \  }}}}t        ||d   z  dk(  d| d|d    d       t        ||d   z  dk(  d| d|d    d       | j                  ||||d   z  |d   ||d   z  |d         } | j                  ddddd	d
      j	                         j                  d||d   |d         }|S )z#Partition windows for NCHW tensors.r   rP  rQ  rR  r-   rS  r   rT  rn   r   r   rU  )r   rL   r   r   r   r   rV  s          r]   window_partition_nchwr    s    JAq!QAA!#xs2QR]^_R`Qaab%cdAA!#wqc1PQ\]^Q_P``a%bc	q!Q+a.(+a.!{1~:M{[\~^Aii1aAq)446;;B;q>S^_`SabGNr_   c           	          |\  }}| j                   d   }| j                  d||d   z  ||d   z  ||d   |d         }|j                  dddddd      j                         j                  d|||      }|S )z*Reverse window partition for NCHW tensors.r-   r   r   rn   rT  r   r   rZ  r[  s          r]   window_reverse_nchwr    s     DAqaARk!n,a;q>.A1kRSnVabcVdeA			!Q1a#..055b!QBAHr_   c           
      h   | j                   \  }}}}t        ||d   z  dk(  d| d|d           t        ||d   z  dk(  d| d|d           | j                  |||d   ||d   z  |d   ||d   z        } | j                  dddddd	      j	                         j                  d
||d   |d         }|S )z Grid partition for NCHW tensors.r   r^  r_  r-   r`  rn   r   r   rT  r   rU  )r   rM   r   r   r   r   rV  s          r]   grid_partition_nchwr    s    JAq!QA	!!WQC/J9UV<.#YZA	!!VA3.I)TU,#XY	q!Yq\1	!#4ilASTDUVAii1aAq)446;;B9Q<QZ[\Q]^GNr_   c           	          |\  }}| j                   d   }| j                  d||d   z  ||d   z  ||d   |d         }|j                  dddddd      j                         j                  d|||      }|S )z(Reverse grid partition for NCHW tensors.r-   r   r   rn   rT  r   r   rZ  rc  s          r]   grid_reverse_nchwr    s     DAqaARil*A1,=q)A,PYZ[P\]A			!Q1a#..055b!QBAHr_   c            	            e Zd ZdZd e       dddfdedededef fd	Zd
e	j                  de	j                  fdZd
e	j                  de	j                  fdZ xZS )PartitionAttention2dzHGrid or Block partition + Attn + FFN.

    '2D' NCHW tensor layout.
    rk  rB   Nr   rl  r  r  c           
      *   ||d}t         |           t        t        |j                        |j
                        }t        |j                        }	|dk(  | _        t        | j                  r|j                  n|j                        | _        t        || j                        }
 ||fi || _        t        ||f|j                   |j"                  |j$                  |
|j&                  |j(                  d|| _        |j,                  rt/        |fd|j,                  i|nt1        j2                         | _        |dkD  rt7        |      nt1        j2                         | _         ||fi || _        t=        d	|t?        ||j@                  z        |	|j(                  d|| _!        |j,                  rt/        |fd|j,                  i|nt1        j2                         | _"        |dkD  rt7        |      | _#        yt1        j2                         | _#        y)
z
        Args:
            dim: Input dimension.
            partition_type: Partition type ('block' or 'grid').
            cfg: Transformer block configuration.
            drop_path: Drop path rate.
        r   r  rk  rn  rP   rB   r  Nri   )$r   r   r   r   rT   rX   r   rR   ro  r$   rL   rM   rp  rh  r  r   r;   rA   r<   rC   rD   r   rP   r   r   r   r  r   r  r  r   rd   r>   r   r  r  rq  s              r]   r   zPartitionAttention2d.__init__  s     /^CNN;N
!#--0	-8'4;O;OUXUbUbc%c4+>+>?*r*


 \\~~#mmmm

 

	 LO??<GGBG`b`k`k`m1:R(9-R[[]*r*
 
c&6&6 67	

 
 LO??<GGBG`b`k`k`m1:R(9-R[[]r_   r   r   c                 0   |j                   dd  }| j                  rt        || j                        }nt	        || j                        }| j                  |      }| j                  rt        || j                  |      }|S t        || j                  |      }|S )Nr   )r   ro  r  rp  r  r   r  r  rs  s       r]   ru  z$PartitionAttention2d._partition_attn  s    7723</43F3FGK-a1D1DEKii,#K1D1DhOA  "+t/B/BHMAr_   c           
      
   || j                  | j                  | j                  | j                  |                        z   }|| j	                  | j                  | j                  | j                  |                        z   }|S rZ   rw  r   s     r]   r   zPartitionAttention2d.forward(  rx  r_   )r`   ra   rb   rc   r7   rd   rh   rg   r   r   r   ru  r   r   r   s   @r]   r  r    s     #*)>)@!1S1S  1S '	1S
 1Sf %,,  %,, r_   r  c                   l     e Zd ZdZd e        e       dddfdedededed	ed
ef fdZddZ	d Z
 xZS )MaxxVitBlockz;MaxVit conv, window partition + FFN , grid partition + FFN.r-   rB   Nr   r   r   r   r   r  c	                 L   ||d}	t         |           |j                  | _        |j                  dk(  rt
        nt        }
 |
||f|||d|	| _        t        d|||d|	}| j                  rt        nt        }|j                  rdn |di || _         |dddi|| _        y)	a^  Initialize MaxxVitBlock.

        Args:
            dim: Input channel dimension.
            dim_out: Output channel dimension.
            stride: Stride for downsampling.
            conv_cfg: Configuration for convolutional blocks.
            transformer_cfg: Configuration for transformer blocks.
            drop_path: Drop path rate.
        r   r   r   r  r  r   r  r  Nrl  gridri   )r   r   rO   	nchw_attnrl   rB  r&  convdictr  rj  rN   r}  r~  )r\   r   r   r   r   r   r  r   r   r   conv_clsr;  partition_layerr   s                r]   r   zMaxxVitBlock.__init__1  s    * /(66$,$7$7:$E=;S'b&hR[b_ab	WwOyWTVW26...FZ"1"?"?$_EcWbEc(NN+Nr_   c                     | j                   %t        t        t        |      | j                          t        t        t        |      | j                         t        t        t
        |      | j                         y r  )r}  r1   r   r   r~  r!  r  r  s     r]   r  zMaxxVitBlock.init_weightsR  sM    ??& 1&A4??SG-f=t~~NGJv6		Br_   c                    | j                  |      }| j                  s|j                  dddd      }| j                  | j                  |      }| j	                  |      }| j                  s|j                  dddd      }|S rM  )r  r  rN  r}  r~  r   s     r]   r   zMaxxVitBlock.forwardX  sp    IIaL~~		!Q1%A??&"ANN1~~		!Q1%Ar_   r  )r`   ra   rb   rc   r6   r7   rd   rg   r   r  r   r   r   s   @r]   r  r  .  sr    E '5'75J5L!OO O 	O
 %O 3O OBCr_   r  c                        e Zd ZdZdd e        e       dddfdededed	ed
ededef fdZdde	ddfdZ
dej                  dej                  fdZ xZS )ParallelMaxxVitBlockzYMaxVit block with parallel cat(window + grid), one FF.

    Experimental timm block.
    r-   r   rB   Nr   r   r   num_convr   r   r  c
                 6   ||	d}
t         |           |j                  dk(  rt        nt        }|dkD  r> |||f|||d|
g}| |||f||d|
g|dz
  z  z  }t        j                  | | _        n |||f|||d|
| _        t        d|||d|
| _	        y)	aa  
        Args:
            dim: Input dimension.
            dim_out: Output dimension.
            stride: Stride for first conv block.
            num_conv: Number of convolution blocks.
            conv_cfg: Convolution block configuration.
            transformer_cfg: Transformer block configuration.
            drop_path: Drop path rate.
        r   r   r-   r  )r  r  r  Nri   )
r   r   rl   rB  r&  r   r  r  rz  r   )r\   r   r   r   r  r   r   r  r   r   r   r  convsr   s                r]   r   zParallelMaxxVitBlock.__init__l  s    , /$,$7$7:$E=;a<c7c6xS\c`bcdEhwYXYVXYZ^fij^jkkEu-DI gff(V_fcefDI.k7[dkhjk	r_   r   r   c                     t        t        t        |      | j                         t        t        t        |      | j
                         y r  )r1   r   r   r   r!  r  r  s     r]   r  z!ParallelMaxxVitBlock.init_weights  s-    G-f=tyyIGJv6		Br_   r   c                     | j                  |      }|j                  dddd      }| j                  |      }|j                  dddd      }|S rM  )r  rN  r   r   s     r]   r   zParallelMaxxVitBlock.forward  sI    IIaLIIaAq!IIaLIIaAq!r_   r  )r`   ra   rb   rc   r6   r7   rd   rg   r   rh   r  r   r   r   r   r   s   @r]   r  r  f  s     '5'75J5L! l l  l 	 l
  l % l 3 l  lDC3 C C %,, r_   r  c                        e Zd ZdZdddd e        e       dddf	ded	ed
ededeeef   dee	ee	   f   dededee
ee
   f   f fdZdej                  dej                  fdZ xZS )MaxxVitStagezEMaxxVit stage consisting of mixed convolution and transformer blocks.r   rT  )   r  r   rB   Nr(  r)  r   depth	feat_sizeblock_typesr   r   r  c           
         |
|d}t         |           d| _        t        ||      }g }t	        |      D ]  \  }}|dk(  r|nd}|dv sJ |dk(  r1|j
                  dk(  rt        nt        }| |||f|||	|   d|gz  }nh|d	k(  r&t        ||      }|t        ||f||||	|   d
|gz  }n=|dk(  r|t        ||f||||	|   d|gz  }n|dk(  r|t        ||f||||	|   d|gz  }|} t        j                  | | _        y)a  
        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            stride: Stride for first block.
            depth: Number of blocks in stage.
            feat_size: Feature map size.
            block_types: Block types ('C' for conv, 'T' for transformer, etc).
            transformer_cfg: Transformer block configuration.
            conv_cfg: Convolution block configuration.
            drop_path: Drop path rate(s).
        r   Fr   r-   )r   r   MPMr   r   r  r   )r   r   r  r  r  )r   r   r   r  r  N)r   r   grad_checkpointingr%   	enumeraterl   rB  r&  rh  r   r  r  r   r  blocks)r\   r(  r)  r   r  r  r  r   r   r  r   r   r   r  itblock_strider  r   r   s                      r]   r   zMaxxVitStage.__init__  s   4 /"'";6k* ,	DAq%&!V6L----Cx,4,?,?:,M=S^8 ( 'l    c-oyI- ( +''l    c< (%$3'l    d/ (%$3'l    FY,	Z mmV,r_   r   r   c                     | j                   r6t        j                  j                         st	        | j
                  |      }|S | j                  |      }|S rZ   )r  r   jitis_scriptingr2   r  r   s     r]   r   zMaxxVitStage.forward  sE    ""599+A+A+Ct{{A.A  AAr_   )r`   ra   rb   rc   r7   r6   rd   r   r   rh   rg   r   r   r   r   r   r   r   s   @r]   r  r    s    O )1255J5L'5'735M-M- M- 	M-
 M- S#XM- sE#J/M- 3M- %M- UDK/0M-^ %,, r_   r  c                        e Zd ZdZ	 	 	 	 	 	 	 	 ddedededededed	ed
ef fdZddeddfdZ	de
j                  de
j                  fdZ xZS )Stemz"Stem layer for feature extraction.Nr(  r)  ro   rw   rG   rR   rT   rX   c                 N   |	|
d}t         |           t        |t        t        f      st        |      }t        t        ||      |      }|d   | _        d| _	        t        ||d   |fd||d|| _         ||d   fi || _        t        |d   |d   |fd||d|| _        y)	ae  
        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            kernel_size: Kernel size for convolutions.
            padding: Padding mode.
            bias: Whether to use bias.
            act_layer: Activation layer.
            norm_layer: Normalization layer.
            norm_eps: Normalization epsilon.
        r   r  r   r   r   )r   rw   rG   r-   N)r   r   r   listtupler$   r   r    r)  r   r!   conv1r  conv2)r\   r(  r)  ro   rw   rG   rR   rT   rX   r   r   r   r5  r   s                r]   r   zStem.__init__  s    0 /'D%=1(G !3J	!JPXYr{"671:{o1V]dholno
#GAJ5"5
"71:wqz;sqZahlsprs
r_   r   r   c                 :    t        t        t        |      |        y r  r=  r  s     r]   r  zStem.init_weights  r>  r_   r   c                 l    | j                  |      }| j                  |      }| j                  |      }|S rZ   )r  r  r  r   s     r]   r   zStem.forward  s.    JJqMJJqMJJqMr_   )rn   rv   FrQ   r   r   NNr  )r`   ra   rb   rc   rd   rh   rf   rg   r   r  r   r   r   r   r   s   @r]   r  r    s    ,  !#+"#t#t #t 	#t
 #t #t #t #t #tJ>3 > > %,, r_   r  c                     | j                   | j                  sJ | S |d   | j                  z  |d   | j                  z  f}t        | ||      } | S )z>Configure window size based on image size and partition ratio.r   r-   )rL   rM   )rL   rM   rK   r   )r  rX  rp  s      r]   cfg_window_sizer  &  sX    
"}}}
a[C$7$77!H[H[9[[N
#>^
LCJr_   kwargsc           	      V   i }i }i }|j                         D ]X  \  }}|j                  d      r|||j                  dd      <   -|j                  d      r|||j                  dd      <   T|||<   Z t        | ft        | j                  fi |t        | j                  fi |d|} | S )z-Overlay keyword arguments onto configuration.transformer_rv   conv_)r   r   )items
startswithr   r   r   )r  r  transformer_kwargsconv_kwargsbase_kwargsr   r   s          r]   _overlay_kwargsr  0  s    KK 1<<'@Aqyy<=\\'"23K		'2./KN  3 3J7IJ55 	C Jr_   c                   z    e Zd ZdZ	 	 	 	 	 	 	 	 d'dedeeeeef   f   dededede	d	e	d
e
f fdZd(dej                  dededdfdZej                   j"                  dee   fd       Zej                   j"                  d)dedeee
f   fd       Zej                   j"                  d*deddfd       Zej                   j"                  dej                  fd       Zd+dedee   ddfdZ	 	 	 	 	 d,dej6                  deeeee   f      dededededeeej6                     eej6                  eej6                     f   f   fdZ	 	 	 d-deeee   f   ded edeed!f   fd"Zdej6                  dej6                  fd#Zd)dej6                  d$edej6                  fd%Z dej6                  dej6                  fd&Z! xZ"S ).r8   z{CoaTNet + MaxVit base model.

    Highly configurable for different block compositions, tensor layouts, pooling types.
    Nr  rX  in_chansnum_classesglobal_pool	drop_ratedrop_path_rater  c
                    t         |           ||	d}t        |      }|
rt        |fi |
}t	        |j
                  |      }|| _        || _        || _        |j                  d   x| _
        | _	        || _        d| _        g | _        t        d||j                  |j                   j"                  |j$                  |j                   j&                  |j                   j(                  |j                   j*                  d|| _        | j,                  j.                  }| xj                  t1        | j,                  j2                  dd      gz  c_        t5        t7        |t        |            D cg c]
  \  }}||z   c}}      }t9        |j                        }t9        |j:                        |k(  sJ t=        ||j:                  d	      }| j,                  j2                  }g }t?        |      D ]  }d}|j                  |   }t5        |D cg c]  }|d
z
  |z  d
z    c}      }|tA        ||f|j:                  |   |jB                  |   |j                   ||||   d|gz  }||z  }|}| xj                  t1        ||d|       gz  c_         tE        jF                  | | _$        tK        tM        |j
                  j(                        |j
                  j*                        }|jN                  rVtE        jP                         | _)        |jN                  | _'        tU        | j                  |f| jN                  |||d|| _+        nH| j                  | _'         || j                  fi || _)        tY        | j                  |f||d|| _+        |jZ                  dv sJ |jZ                  r,t]        tK        | j^                  |jZ                        |        yyc c}}w c c}w )a  
        Args:
            cfg: Model configuration.
            img_size: Input image size.
            in_chans: Number of input channels.
            num_classes: Number of classification classes.
            global_pool: Global pooling type.
            drop_rate: Dropout rate.
            drop_path_rate: Drop path rate.
            **kwargs: Additional keyword arguments to overlay on config.
        r   r   F)r(  r)  rw   rG   rR   rT   rX   r   stem)num_chs	reductionr   T)	stagewiser-   )r  r  r   r   r  r  zstages.r  )hidden_sizerF   r  rT   )rF   r  )rv   r   r   r   r   r  Nri   )0r   r   r$   r  r  r   r  r  r  r   num_featuresr  r  feature_infor  r   r   rw   r   rR   rT   rX   r  r   r  r)  r  ziplenr   r   ranger  rl   r   r  stagesr   r   r   r   r  r   headr   r   r1   _init_weights)r\   r  rX  r  r  r  r  r  r   r   r  r   r   r   r  sr  
num_stagesdprr(  r  stage_strider)  rfinal_norm_layerr   s                            r]   r   zMaxxVit.__init__K  s   0 	/X&!#00C)#*=*=xH& &-0]]2->>DN""' 	
NNLL((ll,,||..\\**	
 	
	 !!d499+<+<RXYZZc(If<M.NOda16OP	'
3::*,,,'

dS""z" 	aALmmA&GINqA,6:NOI|
 jjmNN1- /#a&
 
 
 
F l"FF$w&SZ[\Z]Q^"_!``#	a$ mmV,">#2E2E2P2P#QWZWjWjWsWstDI$'$8$8D!-!! !11%#+ DI %)$5$5D!():):AbADI&!! &#	
 DI "\\\\?? 2 23??KTR m P Os   .O
O 
r   r   r   r   c                     t        |d      r	 |j                  |       y y # t        $ r |j                          Y y w xY w)Nr  r  )hasattrr  	TypeError)r\   r   r   r   s       r]   r  zMaxxVit._init_weights  sD    6>*&##6#2 +  &##%&s   " >>c                     | j                         D ch c]  \  }t        fddD              r c}}S c c}}w )Nc              3   &   K   | ]  }|v  
 y wrZ   ri   ).0nr   s     r]   	<genexpr>z*MaxxVit.no_weight_decay.<locals>.<genexpr>  s     Sa16Ss   )relative_position_bias_tablezrel_pos.mlp)named_parametersany)r\   r   _s    ` r]   no_weight_decayzMaxxVit.no_weight_decay  sK     //1U U!QS#RSS U 	U Us    :coarsec                 $    t        dddg      }|S )Nz^stem)z^stages\.(\d+)N)z^norm)i )r  r  )r  )r\   r  matchers      r]   group_matcherzMaxxVit.group_matcher  s    -/CD
 r_   enablec                 4    | j                   D ]	  }||_         y rZ   )r  r  )r\   r  r  s      r]   set_grad_checkpointingzMaxxVit.set_grad_checkpointing  s     	*A#)A 	*r_   c                 .    | j                   j                  S rZ   )r  fcr[   s    r]   get_classifierzMaxxVit.get_classifier  s    yy||r_   c                 J    || _         | j                  j                  ||       y rZ   )r  r  reset)r\   r  r  s      r]   reset_classifierzMaxxVit.reset_classifier  s    &		[1r_   r   indicesr  
stop_early
output_fmtintermediates_onlyc                    |dv sJ d       g }t        t        | j                        dz   |      \  }}	d}
| j                  |      }|
|v r|j	                  |       t        | j                        }t
        j                  j                         s|s| j                  }n| j                  d|	 }|D ]@  }|
dz  }
 ||      }|
|v s|r|
|k(  r| j                  |      }n|}|j	                  |       B |r|S |
|k(  r| j                  |      }||fS )a   Forward features that returns intermediates.

        Args:
            x: Input image tensor
            indices: Take last n blocks if int, all if None, select matching indices if sequence
            norm: Apply norm layer to compatible intermediates
            stop_early: Stop iterating over blocks when last desired intermediate hit
            output_fmt: Shape of intermediate feature outputs
            intermediates_only: Only return intermediate features
        Returns:

        )NCHWzOutput shape must be NCHW.r-   r   N)	r/   r  r  r  appendr   r  r  r  )r\   r   r  r  r  r  r  intermediatestake_indices	max_indexfeat_idxlast_idxr  stagex_inters                  r]   forward_intermediateszMaxxVit.forward_intermediates  s   * Y&D(DD&"6s4;;7G!7KW"Ui IIaL|#  #t{{#99!!#:[[F[[),F 	.EMHaA<'H0"iilGG$$W-	.   x		!A-r_   
prune_norm
prune_head.c                     t        t        | j                        dz   |      \  }}| j                  d| | _        |rt        j                         | _        |r| j                  dd      | _        |S )z6Prune layers not required for specified intermediates.r-   Nr   rv   )r/   r  r  r   r   r  r  r  )r\   r  r  r  r
  r  s         r]   prune_intermediate_layersz!MaxxVit.prune_intermediate_layers  sb     #7s4;;7G!7KW"Uikk*9-DI--a4DIr_   c                 l    | j                  |      }| j                  |      }| j                  |      }|S rZ   )r  r  r  r   s     r]   forward_featureszMaxxVit.forward_features  s.    IIaLKKNIIaLr_   
pre_logitsc                 N    |r| j                  ||      S | j                  |      S )N)r  )r  )r\   r   r  s      r]   forward_headzMaxxVit.forward_head#  s%    6@tyyzy2RdiiPQlRr_   c                 J    | j                  |      }| j                  |      }|S rZ   )r  r  r   s     r]   r   zMaxxVit.forward&  s'    !!!$a r_   )   rn     r   rB   rB   NNr  F)TrZ   )NFFr  F)r-   FT)#r`   ra   rb   rc   r5   r   rd   r   rh   rg   r   r   r   Moduler  r   r  ignorer   r  rf   r
   r  r  r  r   r  r   r   r  r  r  r  r   r   r   s   @r]   r8   r8   E  s    58#$!$&iSiS CsCx01iS 	iS
 iS iS iS "iS iSV&BII &S &# &t & YYUS U U
 YYD T#s(^   YY*T *T * * YY		  2C 2hsm 2W[ 2 8<$$',4 ||4  eCcN344  	4 
 4  4  !%4  
tELL!5tELL7I)I#JJ	K4 p ./$#	3S	>*  	
 
sCx%,, 5<< Sell S S S %,, r_   r8   r   FTrG   rI   rt   rF   conv_output_biasconv_attn_earlyconv_attn_act_layerconv_norm_layertransformer_shortcut_biastransformer_norm_layertransformer_norm_layer_clrP   rH   rJ   c                 b    t        t        | |dd|||d|	      t        d|||	|||
|            S )as  RW variant configuration for CoAtNet models.

    These models were created and trained before seeing https://github.com/google-research/maxvit

    Common differences for initial timm models:
      - pre-norm layer in MZBConv included an activation after norm
      - mbconv expansion calculated from input instead of output chs
      - mbconv shortcut and final 1x1 conv did not have a bias
      - SE act layer was relu, not silu
      - mbconv uses silu in timm, not gelu
      - expansion in attention block done via output proj, not input proj

    Variable differences (evolved over training initial models):
      - avg pool with kernel_size=2 favoured downsampling (instead of maxpool for coat)
      - SE attention was between conv2 and norm/act
      - default to avg pool for mbconv downsample instead of 1x1 or dw conv
      - transformer block shortcut has no bias
    TFr{   )	rt   rF   rq   rm   rr   rx   r|   rR   rT   )r?   r@   rF   rP   rT   rV   rH   rJ   r   r   r  r6   r7   )rt   rF   r   r!  r"  r#  r$  r%  r&  rP   rH   rJ   s               r]   _rw_coat_cfgr*  ,  sW    @ #(&.&

 .3#-3%#	
 r_   rs   conv_attn_ratior;   c                 `    t        t        | |d||d|      t        d||||	|||
|	            S )a  RW variant configuration for MaxViT models.

    These models were created and trained before seeing https://github.com/google-research/maxvit

    Differences of initial timm models:
      - mbconv expansion calculated from input instead of output chs
      - mbconv shortcut and final 1x1 conv did not have a bias
      - mbconv uses silu in timm, not gelu
      - expansion in attention block done via output proj, not input proj
    Fr{   )rt   rF   rm   rr   r~   rR   rT   )	r?   rF   r;   rL   rP   rT   rV   rH   rJ   r(  r)  )rt   rF   r   r+  r#  r%  r&  rL   r;   rP   rH   rJ   s               r]   _rw_max_cfgr-  e  sS    0 #(&&
 .##-3%#

 r_   rW   r   conv_norm_layer_clrN   c                     t        |      }t        t        d| |d|d   ||      t        d||||d   |||	|
	            S )z=Configuration for experimental ConvNeXt-based MaxxViT models.r   Fr   )rl   rt   rF   rm   rP   rT   rV   r-   )	r?   rF   rL   rN   rP   rT   rV   rH   rJ   r(  )r$   r  r6   r7   )rt   rF   r#  r.  r%  r&  rL   rN   rP   rH   rJ   s              r]   	_next_cfgr0    se     K(K!##A&,
 .#'#A-3%#

 r_   c            	      N    t        t        ddd      t        dddd      	      S )
z0Configuration matching TensorFlow MaxViT models.gMbP?	gelu_tanhsame)rX   rR   rw   r   Frg  )rX   rR   r<   rH   r(  r)  ri   r_   r]   _tf_cfgr4    s6    !

 .!"	
 r_   coatnet_pico_rw)r         rI   r   )r:   r   )r   r   r   r}   )r   r+  coatnet_nano_rw)rn   rT     rn   )rt   r   r+  coatnet_0_rwr   )r   rn   rC  r   )r!  r$  coatnet_1_rw)r   r9  r  r   )rt   r!  r$  coatnet_2_rw)r6  r7  rI      )r   r6  r{   )rt   r"  coatnet_3_rw)r   r   r      )r   r   )rt   r"  rP   coatnet_bn_0_rwr   )rt   r!  r$  r%  coatnet_rmlp_nano_rwr   )r   r+  rH   rJ   coatnet_rmlp_0_rw)rt   rH   coatnet_rmlp_1_rwr   )rF   r!  r$  rH   rJ   coatnet_rmlp_1_rw2)rt   rH   rJ   coatnet_rmlp_2_rw)rt   r"  rP   rH   coatnet_rmlp_3_rwcoatnet_nano_cc)r   r   r   r   rH  )r   r   r   rl   coatnext_nano_rwr   )r   r   r   r   )r   N)rH   rP   	coatnet_0r   r   )r   r   r   r   	coatnet_1	coatnet_2r6  r=  	coatnet_3r   r?  	coatnet_4)r         r   	coatnet_5)r7  rI         rS  maxvit_pico_rw)r:   r   r6  r7  )r   r   r   r   )r  r  r  r  )   r:   )r   r   rl   r   maxvit_nano_rw)r-   r   rn   r-   maxvit_tiny_rwmaxvit_tiny_pm)r  r  r  r  maxvit_rmlp_pico_rw)rH   maxvit_rmlp_nano_rwmaxvit_rmlp_tiny_rwmaxvit_rmlp_small_rwmaxvit_rmlp_base_rw)r   r   rl   r   r   maxxvit_rmlp_nano_rw)r   r   rl   r   r   maxxvit_rmlp_tiny_rwmaxxvit_rmlp_small_rw)0   r   maxxvitv2_nano_rw)rN   rH   maxxvitv2_rmlp_base_rw)r   r9  rO  r   )rN   maxxvitv2_rmlp_large_rw)   i@  i  rR  )r   r9     r   )P   re  rR  maxvit_tiny_tf)r   r   rl   r   r   r   maxvit_small_tfmaxvit_base_tfmaxvit_large_tfmaxvit_xlarge_tf
state_dictmodelc                 r   |j                         }i }| j                         D ]  \  }}|j                  d      r|j                  |dd       }|j                  |j
                  j                  k7  s|j                  d   |j                  d   k7  r,t        ||j                  |j
                  j                        }||v rn|j                  ||   j                  k7  rR|j                         ||   j                         k(  r.|j                  dv sJ |j                  ||   j                        }|||<    |S )z/Filter checkpoint state dict for compatibility.r  Nir   r-   )new_window_sizenew_bias_shape)r   rT  )rm  r  endswithget_submoduler   r  rL   r,   ndimnumelr   )rm  rn  model_state_dictout_dictr   r   ms          r]   checkpoint_filter_fnry  >  s)   '')H  " 1::45##AdsG,Aww!88>>>!--PQBRVWVcVcdeVfBf-$%MM#$#A#A#G#G   QVV/?/B/G/G%GAGGIYijkYlYrYrYtLt66V###		*1-334A Or_   variantcfg_variant
pretrainedc                     |.| t         v r| }n#dj                  | j                  d      dd       }t        t        | |ft         |   t        d      t        d|S )zCreate a MaxxVit model variant.Nr  r   T)flatten_sequential)	model_cfgfeature_cfgpretrained_filter_fn)
model_cfgsjoinsplitr.   r8   r  ry  )rz  r{  r|  r  s       r]   _create_maxxvitr  T  si    j !K((7==#5cr#:;K*[)D11	
  r_   urlc                 $    | dddddddddd	d
d|S )z$Create a default configuration dict.r  )rn   r  r  )rC  rC  ffffff?bicubic)      ?r  r  z
stem.conv1zhead.fcTz
apache-2.0)r  r  
input_size	pool_sizecrop_pctinterpolationmeanr   
first_conv
classifierfixed_input_sizelicenseri   )r  r  s     r]   _cfgr  c  s7     4}SY9")  $* r_   zcoatnet_pico_rw_224.untrained)r  zcoatnet_nano_rw_224.sw_in1kztimm/zyhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_nano_rw_224_sw-f53093b4.pthg?)	hf_hub_idr  r  zcoatnet_0_rw_224.sw_in1kzvhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_0_rw_224_sw-a6439706.pth)r  r  zcoatnet_1_rw_224.sw_in1kzvhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_1_rw_224_sw-5cae1ea8.pthz!coatnet_2_rw_224.sw_in12k_ft_in1k)r  z'coatnet_rmlp_1_rw2_224.sw_in12k_ft_in1kz&coatnet_rmlp_2_rw_224.sw_in12k_ft_in1kz&coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k)rn   r   r   )rO  rO  g      ?squash)r  r  r  r  	crop_modezcoatnet_bn_0_rw_224.sw_in1kzyhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_bn_0_rw_224_sw-c228e218.pthr  )r  r  r  r   r  z coatnet_rmlp_nano_rw_224.sw_in1kz~https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_rmlp_nano_rw_224_sw-bd1d51b3.pthzcoatnet_rmlp_0_rw_224.untrainedzcoatnet_rmlp_1_rw_224.sw_in1kz{https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_rmlp_1_rw_224_sw-9051e6c3.pthzcoatnet_rmlp_2_rw_224.sw_in1kz{https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_rmlp_2_rw_224_sw-5ccfac55.pthzcoatnet_rmlp_3_rw_224.untrainedzcoatnet_nano_cc_224.untrainedzcoatnext_nano_rw_224.sw_in1kzzhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnext_nano_rw_224_ad-22cb71c2.pthzcoatnet_2_rw_224.sw_in12ki-.  )r  r  zcoatnet_3_rw_224.sw_in12kzcoatnet_rmlp_1_rw2_224.sw_in12kzcoatnet_rmlp_2_rw_224.sw_in12kzcoatnet_0_224.untrainedzcoatnet_1_224.untrainedzcoatnet_2_224.untrainedzcoatnet_3_224.untrainedzcoatnet_4_224.untrainedzcoatnet_5_224.untrainedzmaxvit_pico_rw_256.untrained)rn   r7  r7  )   r  )r  r  r  zmaxvit_nano_rw_256.sw_in1kzxhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_nano_rw_256_sw-fb127241.pth)r  r  r  r  zmaxvit_tiny_rw_224.sw_in1kzxhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_tiny_rw_224_sw-7d0dffeb.pthzmaxvit_tiny_rw_256.untrainedzmaxvit_tiny_pm_256.untrainedzmaxvit_rmlp_pico_rw_256.sw_in1kz}https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_rmlp_pico_rw_256_sw-8d82f2c6.pthzmaxvit_rmlp_nano_rw_256.sw_in1kz}https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_rmlp_nano_rw_256_sw-c17bb0d6.pthzmaxvit_rmlp_tiny_rw_256.sw_in1kz}https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_rmlp_tiny_rw_256_sw-bbef0ff5.pthz maxvit_rmlp_small_rw_224.sw_in1kz~https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_rmlp_small_rw_224_sw-6ef0ae4f.pthz"maxvit_rmlp_small_rw_256.untrainedz(maxvit_rmlp_base_rw_224.sw_in12k_ft_in1kz(maxvit_rmlp_base_rw_384.sw_in12k_ft_in1kz maxvit_rmlp_base_rw_224.sw_in12kz maxxvit_rmlp_nano_rw_256.sw_in1kz~https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxxvit_rmlp_nano_rw_256_sw-0325d459.pthz"maxxvit_rmlp_tiny_rw_256.untrainedz!maxxvit_rmlp_small_rw_256.sw_in1kzhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxxvit_rmlp_small_rw_256_sw-37e217ff.pthzmaxxvitv2_nano_rw_256.sw_in1k)r  r  r  z+maxxvitv2_rmlp_base_rw_224.sw_in12k_ft_in1kz+maxxvitv2_rmlp_base_rw_384.sw_in12k_ft_in1kz%maxxvitv2_rmlp_large_rw_224.untrainedz#maxxvitv2_rmlp_base_rw_224.sw_in12kzmaxvit_tiny_tf_224.in1k)r  r  r   zmaxvit_tiny_tf_384.in1kzmaxvit_tiny_tf_512.in1k)rn   rI   rI   )rf  rf  zmaxvit_small_tf_224.in1kzmaxvit_small_tf_384.in1kzmaxvit_small_tf_512.in1kzmaxvit_base_tf_224.in1kzmaxvit_base_tf_384.in1kzmaxvit_base_tf_512.in1kzmaxvit_large_tf_224.in1kzmaxvit_large_tf_384.in1kzmaxvit_large_tf_512.in1kzmaxvit_base_tf_224.in21kiSU  z maxvit_base_tf_384.in21k_ft_in1kz maxvit_base_tf_512.in21k_ft_in1kzmaxvit_large_tf_224.in21kz!maxvit_large_tf_384.in21k_ft_in1kz!maxvit_large_tf_512.in21k_ft_in1k)r  r  r  r  zmaxvit_xlarge_tf_224.in21kz"maxvit_xlarge_tf_384.in21k_ft_in1kz"maxvit_xlarge_tf_512.in21k_ft_in1kc                     t        dd| i|S )z)CoatNet Pico model with RW configuration.r|  )coatnet_pico_rw_224r  r|  r  s     r]   r  r  J	       RZR6RRr_   c                     t        dd| i|S )z)CoatNet Nano model with RW configuration.r|  )coatnet_nano_rw_224r  r  s     r]   r  r  P	  r  r_   c                     t        dd| i|S )z&CoatNet-0 model with RW configuration.r|  )coatnet_0_rw_224r  r  s     r]   r  r  V	       O*OOOr_   c                     t        dd| i|S )z&CoatNet-1 model with RW configuration.r|  )coatnet_1_rw_224r  r  s     r]   r  r  \	  r  r_   c                     t        dd| i|S )z&CoatNet-2 model with RW configuration.r|  )coatnet_2_rw_224r  r  s     r]   r  r  b	  r  r_   c                     t        dd| i|S )z&CoatNet-3 model with RW configuration.r|  )coatnet_3_rw_224r  r  s     r]   r  r  h	  r  r_   c                     t        dd| i|S )z4CoatNet-0 model with BatchNorm and RW configuration.r|  )coatnet_bn_0_rw_224r  r  s     r]   r  r  n	  r  r_   c                     t        dd| i|S )z.CoatNet Nano model with Relative Position MLP.r|  )coatnet_rmlp_nano_rw_224r  r  s     r]   r  r  t	       W*WPVWWr_   c                     t        dd| i|S )z+CoatNet-0 model with Relative Position MLP.r|  )coatnet_rmlp_0_rw_224r  r  s     r]   r  r  z	       TzTVTTr_   c                     t        dd| i|S )z+CoatNet-1 model with Relative Position MLP.r|  )coatnet_rmlp_1_rw_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z.CoatNet-1 model with Relative Position MLP v2.r|  )coatnet_rmlp_1_rw2_224r  r  s     r]   r  r  	  s     U
UfUUr_   c                     t        dd| i|S )z+CoatNet-2 model with Relative Position MLP.r|  )coatnet_rmlp_2_rw_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z6CoatNet-2 model with Relative Position MLP at 384x384.r|  )coatnet_rmlp_2_rw_384r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z+CoatNet-3 model with Relative Position MLP.r|  )coatnet_rmlp_3_rw_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z(CoatNet Nano model with ConvNeXt blocks.r|  )coatnet_nano_cc_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z*CoAtNeXt Nano model with RW configuration.r|  )coatnext_nano_rw_224r  r  s     r]   r  r  	  s     SjSFSSr_   c                     t        dd| i|S )zCoatNet-0 model.r|  )coatnet_0_224r  r  s     r]   r  r  	       LzLVLLr_   c                     t        dd| i|S )zCoatNet-1 model.r|  )coatnet_1_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )zCoatNet-2 model.r|  )coatnet_2_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )zCoatNet-3 model.r|  )coatnet_3_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )zCoatNet-4 model.r|  )coatnet_4_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )zCoatNet-5 model.r|  )coatnet_5_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z(MaxViT Pico model with RW configuration.r|  )maxvit_pico_rw_256r  r  s     r]   r  r  	       QJQ&QQr_   c                     t        dd| i|S )z(MaxViT Nano model with RW configuration.r|  )maxvit_nano_rw_256r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z(MaxViT Tiny model with RW configuration.r|  )maxvit_tiny_rw_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z3MaxViT Tiny model with RW configuration at 256x256.r|  )maxvit_tiny_rw_256r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z3MaxViT Relative Position MLP Pico RW 256x256 model.r|  )maxvit_rmlp_pico_rw_256r  r  s     r]   r  r  	       VVvVVr_   c                     t        dd| i|S )z3MaxViT Relative Position MLP Nano RW 256x256 model.r|  )maxvit_rmlp_nano_rw_256r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z3MaxViT Relative Position MLP Tiny RW 256x256 model.r|  )maxvit_rmlp_tiny_rw_256r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z4MaxViT Relative Position MLP Small RW 224x224 model.r|  )maxvit_rmlp_small_rw_224r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z9MaxViT Small model with Relative Position MLP at 256x256.r|  )maxvit_rmlp_small_rw_256r  r  s     r]   r  r  	  r  r_   c                     t        dd| i|S )z-MaxViT Base model with Relative Position MLP.r|  )maxvit_rmlp_base_rw_224r  r  s     r]   r  r  
  r  r_   c                     t        dd| i|S )z8MaxViT Base model with Relative Position MLP at 384x384.r|  )maxvit_rmlp_base_rw_384r  r  s     r]   r  r  

  r  r_   c                     t        dd| i|S )z'MaxViT Tiny model with parallel blocks.r|  )maxvit_tiny_pm_256r  r  s     r]   r  r  
  r  r_   c                     t        dd| i|S )z4MaxxViT Relative Position MLP Nano RW 256x256 model.r|  )maxxvit_rmlp_nano_rw_256r  r  s     r]   r  r  
  r  r_   c                     t        dd| i|S )z.MaxxViT Tiny model with Relative Position MLP.r|  )maxxvit_rmlp_tiny_rw_256r  r  s     r]   r  r  
  r  r_   c                     t        dd| i|S )z/MaxxViT Small model with Relative Position MLP.r|  )maxxvit_rmlp_small_rw_256r  r  s     r]   r  r  "
  s     X:XQWXXr_   c                     t        dd| i|S )zMaxxViT-V2 Nano model.r|  )maxxvitv2_nano_rw_256r  r  s     r]   r  r  (
  r  r_   c                     t        dd| i|S )z1MaxxViT-V2 Base model with Relative Position MLP.r|  )maxxvitv2_rmlp_base_rw_224r  r  s     r]   r  r  .
       YJYRXYYr_   c                     t        dd| i|S )z<MaxxViT-V2 Base model with Relative Position MLP at 384x384.r|  )maxxvitv2_rmlp_base_rw_384r  r  s     r]   r  r  4
  r  r_   c                     t        dd| i|S )z2MaxxViT-V2 Large model with Relative Position MLP.r|  )maxxvitv2_rmlp_large_rw_224r  r  s     r]   r  r  :
  s     ZZZSYZZr_   c                     t        dd| i|S )z"MaxViT Tiny model from TensorFlow.r|  )maxvit_tiny_tf_224rh  r  r  s     r]   r  r  @
       cjc\bccr_   c                     t        dd| i|S )z-MaxViT Tiny model from TensorFlow at 384x384.r|  )maxvit_tiny_tf_384rh  r  r  s     r]   r  r  F
  r  r_   c                     t        dd| i|S )z-MaxViT Tiny model from TensorFlow at 512x512.r|  )maxvit_tiny_tf_512rh  r  r  s     r]   r  r  L
  r  r_   c                     t        dd| i|S )z#MaxViT Small model from TensorFlow.r|  )maxvit_small_tf_224ri  r  r  s     r]   r  r  R
       ePZe^deer_   c                     t        dd| i|S )z.MaxViT Small model from TensorFlow at 384x384.r|  )maxvit_small_tf_384ri  r  r  s     r]   r   r   X
  r  r_   c                     t        dd| i|S )z.MaxViT Small model from TensorFlow at 512x512.r|  )maxvit_small_tf_512ri  r  r  s     r]   r  r  ^
  r  r_   c                     t        dd| i|S )z"MaxViT Base model from TensorFlow.r|  )maxvit_base_tf_224rj  r  r  s     r]   r  r  d
  r  r_   c                     t        dd| i|S )z-MaxViT Base model from TensorFlow at 384x384.r|  )maxvit_base_tf_384rj  r  r  s     r]   r  r  j
  r  r_   c                     t        dd| i|S )z-MaxViT Base model from TensorFlow at 512x512.r|  )maxvit_base_tf_512rj  r  r  s     r]   r  r  p
  r  r_   c                     t        dd| i|S )z#MaxViT Large model from TensorFlow.r|  )maxvit_large_tf_224rk  r  r  s     r]   r
  r
  v
  r  r_   c                     t        dd| i|S )z.MaxViT Large model from TensorFlow at 384x384.r|  )maxvit_large_tf_384rk  r  r  s     r]   r  r  |
  r  r_   c                     t        dd| i|S )z.MaxViT Large model from TensorFlow at 512x512.r|  )maxvit_large_tf_512rk  r  r  s     r]   r  r  
  r  r_   c                     t        dd| i|S )z$MaxViT XLarge model from TensorFlow.r|  )maxvit_xlarge_tf_224rl  r  r  s     r]   r  r  
       gR\g`fggr_   c                     t        dd| i|S )z/MaxViT XLarge model from TensorFlow at 384x384.r|  )maxvit_xlarge_tf_384rl  r  r  s     r]   r  r  
  r  r_   c                     t        dd| i|S )z/MaxViT XLarge model from TensorFlow at 512x512.r|  )maxvit_xlarge_tf_512rl  r  r  s     r]   r  r  
  r  r_   r  )r   rE   FFrelurv   TrS   rU   NrG   rI   )rs   rE   Fg      ?rv   rS   rU   Nr:   NrG   rI   )rs   rE   rS   rU   rS   rU   NFrW   r   rI   ri   )NFr  )rc   r  collectionsr   dataclassesr   r   r   	functoolsr   typingr   r	   r
   r   r   r   r   r   r   r   	torch.jitr   	timm.datar   r   timm.layersr   r   r   r   r   r   r   r   r   r   r   r   r    r!   r"   r#   r$   r%   r&   r'   r(   r)   r*   r+   r,   _builderr.   	_featuresr/   _features_fxr0   _manipulater1   r2   	_registryr3   r4   __all__r7   r6   r5   r  r   r   r   rh   r   r   r!  rd   r$  r&  rB  r   rW  r\  ra  rd  rh  rj  rz  r  r  r  r  r  r  r  r  r  r  r  r8   rf   rg   r*  r-  r0  r4  r  r  ry  r  r  default_cfgsr  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r   r  r  r  r  r
  r  r  r  r  r  ri   r_   r]   <module>r%     s   "H  # 1 1  I I I    A      6 + + 3 4 <
N 2 2 2B !P !P !PH 
! 
! 
!Q")) QhR")) Rj/299 /d0bii 0s 0C 0 02N Nb,ryy , ,S ,$ ,0&8C= &C &C &l")) l^dBII dN 49   ELL tCy DQTI Z_ZfZf  ell tCy U\\  %,, 49 S	 V[VbVb  	. 	U38_ 	QYZbQc 	D299 DNT TnU\\ S	 ell   DI QUVYQZ _d_k_k  5<< DI %,,  u|| S	 TRUY [`[g[g  K299 K\5299 5p1299 1hW299 Wt/299 /d. %S/ Nc  s z *dbii dP "!& %#)!*.&3)4'+"666 6 	6
 !6 6 $(6 !$6 $'6 e_6 6 6 
#s(^6t  !&!'!&3)415'+"--- - 	-
 - !$- $'- eCHo.- - e_- - - 
#s(^-b  ,"-&3)415#9=!$$$ $  	$
 !$$ $'$ eCHo.$ $ 5%u"556$ $ $ 
#s(^$Nc3h "  n % ! 
	n  	%	 ! 
		n*  %  &+
	+n<  	%	  &+
		=nP  	'	  &
		Qnd  	'	  &
		en|  
%
  &+#0	
	
}nR $ 
%
 ! 	
	
Snh ! % 
	inz ! %  &+
	{nR " 	%	 
		Snf ! 
'
  &	
	
gn| ! 
'
  &	
	
}nT  %5	
 .Unb   	%		
 $
	cnz %	{nF %	GnR '	Sn^ '	_nj '	knv (	wnF  $	
 -GnT  %	
 -Unb  %	
 -cnp  %	
 -qn@ # $	
 5
)AnN # %	
 5
)On\ # %	
 5
)]nj $ 	%		
 
	kn~ # 	%	 
	nT $ % +Und $ %	
 +enr % %	
 +snB	 ! 
%
 

C	nX	 & '	
 
Y	nj	 ' 	'	 
	k	nB
  % )C
nT
  % )U
nf
  % )g
nx
  ' )y
nJ   ' )Kn
bT#u||*;%< RYY SWX[]b]i]iXiSj ,S x} QU il qx 	c 	# 	$sCx. 	 % X&#Tb\X& "4 H$X&  E!FX&  E!X&  (*!X&* .t0+X&. -d//X&2 -d Hsh/X3X&< "4 H"(<	$=X&F ' M)GX&N &t|OX&P $T J&KQX&V $T J&KWX&\ &t|]X&^ $Tb\_X&` #D I%aX&l  "mX&r  "sX&x &t(yX&~ %d'X&H t|IX&J t|KX&L t|MX&N t|OX&P t|QX&R t|SX&X #DRMU[$\YX&Z !$ G F#4[X&b !$ G#HcX&h #D F%4iX&n #DRMU[$\oX&t &t L F(4uX&| &t L F(4}X&D &t L F(4EX&L ' M)MX&V )$ F+4WX&` /1aX&f / Hsh1XgX&p ')qX&| ' M F)4}X&D )$2-[a*bEX&F ( N F*4GX&R $T F&4SX&X 244YX&\ 24 Hsh4X]X&b ,Tb\cX&f *4,gX&p t"(< >qX&v t Hsh XwX&| t Hsh X}X&B "(<!>CX&H  Hsh!XIX&N  Hsh!XOX&T t"(< >UX&Z t Hsh X[X&` t Hsh XaX&f "(<!>gX&l  Hsh!XmX&r  Hsh!XsX&z !{X&@ ' Hsh)XAX&F ' Hsh)XGX&L  "MX&R ( Hsh*XSX&X ( 3(*DYX&^ !$#_X&d )$ Hsh+XeX&j )$ Hsh+XkX& Xv SD SC SG S S
 SD SC SG S S
 P P P P P
 P P P P P
 P P P P P
 P P P P P
 SD SC SG S S
 X X X X X
 Ud Uc Ug U U
 Ud Uc Ug U U
 Vt Vs Vw V V
 Ud Uc Ug U U
 Ud Uc Ug U U
 Ud Uc Ug U U
 SD SC SG S S
 TT TS TW T T
 Md Mc Mg M M
 Md Mc Mg M M
 Md Mc Mg M M
 Md Mc Mg M M
 Md Mc Mg M M
 Md Mc Mg M M
 R4 R3 R7 R R
 R4 R3 R7 R R
 R4 R3 R7 R R
 R4 R3 R7 R R
 W W W W W
 W W W W W
 W W W W W
 X X X X X
 X X X X X
 W W W W W
 W W W W W
 R4 R3 R7 R R
 X X X X X
 X X X X X
 Y$ Y# Y' Y Y
 Ud Uc Ug U U
 Z4 Z3 Z7 Z Z
 Z4 Z3 Z7 Z Z
 [D [C [G [ [
 d4 d3 d7 d d
 d4 d3 d7 d d
 d4 d3 d7 d d
 fD fC fG f f
 fD fC fG f f
 fD fC fG f f
 d4 d3 d7 d d
 d4 d3 d7 d d
 d4 d3 d7 d d
 fD fC fG f f
 fD fC fG f f
 fD fC fG f f
 hT hS hW h h
 hT hS hW h h
 hT hS hW h hr_   