
    ^j                     "   d dl Z d dlmZ d dlmZmZmZmZmZ d dl	Z	d dl
mZ d dlmZmZ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 ddlmZ ddlmZ dd	lm Z  dd
l!m"Z"m#Z# dgZ$d Z% G d dejL                        Z' G d dejL                        Z(ejR                  ddddfde*de*deejL                     de+de+dejX                  fdZ- G d dejL                        Z. G d dejL                        Z/ G d dejL                        Z0 G d  d!ejL                        Z1 G d" d#ejL                        Z2 G d$ d%ejL                        Z3 G d& d'ejL                        Z4 G d( d)ejL                        Z5 G d* d+ejL                        Z6 G d, dejL                        Z7dcd-Z8 e#i d. e8d/0      d1 e8d/0      d2 e8d/0      d3 e8d/0      d4 e8d/0      d5 e8d/0      d6 e8d/d78      d9 e8d/0      d: e8d/0      d; e8d/0      d< e8d/0      d= e8d/0      d> e8d/0      d? e8d/d78      d@ e8dAdBd7dCdDdEdFG      dH e8dIdJd7dCdDdEdFG      dK e8dLdMd7dCdDdEdFG       e8d/dNdCdDdEdFO       e8d/d7dCdDdEdFO       e8d/d7dPeedQdRdFS       e8d/d7dPeedQdRdFS      dT      Z9dU Z:dddVZ;e"dddW       Z<e"dddX       Z=e"dddY       Z>e"dddZ       Z?e"ddd[       Z@e"ddd\       ZAe"ddd]       ZBe"ddd^       ZCe"ddd_       ZDe"ddd`       ZEe"ddda       ZFe"dddb       ZGy)e    N)partial)ListOptionalTupleTypeUnion)IMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STDOPENAI_CLIP_MEANOPENAI_CLIP_STD)	DropPathcalculate_drop_path_ratestrunc_normal_create_conv2dConvNormActSqueezeExciteuse_fused_attnClassifierHeadLayerNorm2d   )build_model_with_cfg)feature_take_indices)checkpoint_seq)register_modelgenerate_default_cfgsFastVitc                 &    | sy|| z  dk(  sJ || z  S )Nr   r    )
group_sizechannelss     ^/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/fastvit.py
num_groupsr"   "   s(     *$))):%%    c                       e Zd ZdZddddddddej
                  ddfdeded	ed
ededededededededeej                     ddf fdZ
dej                  dej                  fdZd Zdeej                  ej                  f   fdZdeej$                  ej&                  f   deej                  ej                  f   fdZ xZS )MobileOneBlocka#  MobileOne building block.

    This block has a multi-branched architecture at train-time
    and plain-CNN style architecture at inference time
    For more details, please refer to our paper:
    `An Improved One millisecond Mobile Backbone` -
    https://arxiv.org/pdf/2206.04040.pdf
    r   r   FTNin_chsout_chskernel_sizestridedilationr   inference_modeuse_seuse_actuse_scale_branchnum_conv_branches	act_layerreturnc                    ||d}t         |           || _        t        ||      | _        || _        || _        || _        || _        || _	        || _
        |rt        |fddi|nt        j                         | _        |r#t        ||f|||| j                  dd|| _        nd| _        ||k(  r|dk(  rt        j"                  dd|i|nd| _        |dkD  rtt        j&                  t)        | j                        D cg c]<  }t+        | j                  | j                  f|| j
                  | j                  d	d
|> c}      | _        nd| _        d| _        |dkD  rA|
r?t+        | j                  | j                  fd| j
                  | j                  d	d
|| _        |	r |       | _        yt        j                         | _        yc c}w )a  Construct a MobileOneBlock module.

        Args:
            in_chs: Number of channels in the input.
            out_chs: Number of channels produced by the block.
            kernel_size: Size of the convolution kernel.
            stride: Stride size.
            dilation: Kernel dilation factor.
            group_size: Convolution group size.
            inference_mode: If True, instantiates model in inference mode.
            use_se: Whether to use SE-ReLU activations.
            use_act: Whether to use activation. Default: ``True``
            use_scale_branch: Whether to use scale branch. Default: ``True``
            num_conv_branches: Number of linear conv branches.
        devicedtype
rd_divisorr   Tr(   r)   r*   groupsbiasNnum_featuresr   Fr(   r)   r8   	apply_actr   )super__init__r+   r"   r8   r)   r*   r(   r&   r'   r/   r   nnIdentityser   reparam_convBatchNorm2didentity
ModuleListranger   conv_kxk
conv_scaleact)selfr&   r'   r(   r)   r*   r   r+   r,   r-   r.   r/   r0   r4   r5   dd_	__class__s                    r!   r>   zMobileOneBlock.__init__5   s   @ /, V4 &!2 AG-<A<<BKKM -	! (!{{	! 	!D !%D f$1 9F9b9 M !1$ "  %T%;%;<
/    %0#{{#{{"' 
/ 
! !% #DOQ#3"-KKLL# !";;;;## # #*9;r{{}7
/s   ?AG
xc                    | j                   /| j                  | j                  | j                  |                  S d}| j                  | j                  |      }d}| j                  | j	                  |      }||z   }| j
                  | j
                  D ]  }| ||      z  } | j                  | j                  |            S )zApply forward pass.r   )rB   rI   rA   rD   rH   rG   )rJ   rN   identity_out	scale_outoutrcs         r!   forwardzMobileOneBlock.forward   s     (88DGGD$5$5a$89:: ==$==+L 	??&*I ,&==$mm r!u xx%%r#   c           	      <   | j                   y| j                         \  }}t        | j                  | j                  | j
                  | j                  | j                  | j                  d      | _         || j                   j                  _
        || j                   j                  _
        | j                         D ]  \  }}d|v r|j                           | j                  d       | j                  d       t        | d      r| j                  d       d| _        y)a  Following works like `RepVGG: Making VGG-style ConvNets Great Again` -
        https://arxiv.org/pdf/2101.03697.pdf. We re-parameterize multi-branched
        architecture used at training time to obtain a plain CNN-like structure
        for inference.
        NT)in_channelsout_channelsr(   r)   r*   r8   r9   rB   rG   rH   rD   )rB   _get_kernel_biasr   r&   r'   r(   r)   r*   r8   weightdatar9   named_parametersdetach___delattr__hasattrr+   )rJ   kernelr9   nameparas        r!   reparameterizezMobileOneBlock.reparameterize   s     (,,.)((;;]];;
 )/  %&*# //1 	JD$%LLN	
 	$&4$Z("r#   c                    d}d}| j                   [| j                  | j                         \  }}| j                  dz  }t        j                  j
                  j                  |||||g      }d}d}| j                  | j                  | j                        \  }}d}d}| j                  Et        | j                        D ]-  }| j                  | j                  |         \  }	}
||	z  }||
z  }/ ||z   |z   }||z   |z   }||fS )zMethod to obtain re-parameterized kernel and bias.
        Reference: https://github.com/DingXiaoH/RepVGG/blob/main/repvgg.py#L83

        Returns:
            Tuple of (kernel, bias) after fusing branches.
        r      )rH   _fuse_bn_tensorr(   torchr?   
functionalpadrD   rG   rF   r/   )rJ   kernel_scale
bias_scalerh   kernel_identitybias_identitykernel_conv	bias_convix_kernel_biaskernel_final
bias_finals                r!   rX   zMobileOneBlock._get_kernel_bias   s"    
??&'+';';DOO'L$L*""a'C 88..22<#sCQTAUVL ==$-1-A-A$---P*O] 	==$D223 #!%!5!5dmmB6G!Hw&U"	#
 #\1OC+m;
Z''r#   branchc                    t        |t              r|j                  j                  }|j                  j
                  }|j                  j                  }|j                  j                  }|j                  j                  }|j                  j                  }n2t        |t        j                        sJ t        | d      s| j                  | j                  z  }t        j                  | j                  || j                   | j                   f|j                  j"                  |j                  j$                        }	t'        | j                        D ](  }
d|	|
|
|z  | j                   dz  | j                   dz  f<   * |	| _        | j(                  }|j
                  }|j                  }|j                  }|j                  }|j                  }||z   j+                         }||z  j-                  dddd      }||z  |||z  |z  z
  fS )a  Method to fuse batchnorm layer with preceding conv layer.
        Reference: https://github.com/DingXiaoH/RepVGG/blob/main/repvgg.py#L95

        Args:
            branch: Sequence of ops to be fused.

        Returns:
            Tuple of (kernel, bias) after fusing batchnorm.
        	id_tensorr5   r4   r   rd   )
isinstancer   convrY   bnrunning_meanrunning_varr9   epsr?   rC   r^   r&   r8   rf   zerosr(   r5   r4   rF   rv   sqrtreshape)rJ   rt   r_   r|   r}   gammabetar~   	input_dimkernel_valueistdts                r!   re   zMobileOneBlock._fuse_bn_tensor   s    fk*[[''F!9911L ))//KII$$E99>>D))--Cfbnn5554- KK4;;6	${{[[)T-=-=t?O?OP ----!==// 
 t{{+ A  !1y=$*:*:a*?AQAQUVAVV ".^^F!..L ,,KMME;;D**CS &&(S[!!"aA.z4,"6"<<<<r#   )__name__
__module____qualname____doc__r?   GELUintboolr   Moduler>   rf   TensorrT   rb   r   rX   r   
SequentialrC   re   __classcell__rM   s   @r!   r%   r%   +   sF    #(  %)%&)+`=`= `= 	`=
 `= `= `= !`= `= `= #`=  #`= BII`=  
!`=D& &%,, &2!#F!(%ell(B"C !(F*="--78*= 
u||U\\)	**=r#   r%   c                   p    e Zd ZdZ	 	 	 	 	 	 ddedededededee   d	ed
eej                     deddf fdZ	de
j                  de
j                  fdZdee
j                  e
j                  f   fdZddZedej"                  dej$                  dee
j                  e
j                  f   fd       Z xZS )ReparamLargeKernelConvzBuilding Block of RepLKNet

    This class defines overparameterized large kernel conv block
    introduced in `RepLKNet <https://arxiv.org/abs/2203.06717>`_

    Reference: https://github.com/DingXiaoH/RepLKNet-pytorch
    Nr&   r'   r(   r)   r   small_kernelr,   r0   r+   r1   c           	      V   |
|d}t         |           || _        t        ||      | _        || _        || _        || _        || _        |	r#t        ||f||d| j                  dd|| _
        nkd| _
        t        ||f|| j                  | j                  dd|| _        |7||k  sJ d       t        ||f|| j                  | j                  dd|| _        |rt        |fd	d
i|nt        j                          | _        | |       | _        yt        j                          | _        y)a!  Construct a ReparamLargeKernelConv module.

        Args:
            in_chs: Number of input channels.
            out_chs: Number of output channels.
            kernel_size: Kernel size of the large kernel conv branch.
            stride: Stride size. Default: 1
            group_size: Group size. Default: 1
            small_kernel: Kernel size of small kernel conv branch.
            act_layer: Activation module. Default: ``nn.GELU``
            inference_mode: If True, instantiates model in inference mode. Default: ``False``
        r3   r   Tr7   NFr;   zDThe kernel size for re-param cannot be larger than the large kernel!rd_ratiog      ?)r=   r>   r)   r"   r8   r&   r'   r(   r   r   rB   r   
large_conv
small_convr   r?   r@   rA   rI   )rJ   r&   r'   r(   r)   r   r   r,   r0   r+   r4   r5   rK   rM   s                r!   r>   zReparamLargeKernelConv.__init__,  s_   4 / V4&( -	! ({{	! 	!D !%D) ({{{{ DO ' K/ZYZ/"-# !-;;;;## # BH-=$="=R[[]"+"79;R[[]r#   rN   c                     | j                   | j                  |      }n1| j                  |      }| j                  || j                  |      z   }| j                  |      }| j	                  |      }|S N)rB   r   r   rA   rI   )rJ   rN   rR   s      r!   rT   zReparamLargeKernelConv.forwardv  sh    (##A&C//!$C*DOOA..ggclhhsm
r#   c                    | j                  | j                  j                  | j                  j                        \  }}t	        | d      r| j                  | j
                  j                  | j
                  j                        \  }}||z  }|t        j                  j                  || j                  | j                  z
  dz  gdz        z  }||fS )zMethod to obtain re-parameterized kernel and bias.
        Reference: https://github.com/DingXiaoH/RepLKNet-pytorch

        Returns:
            Tuple of (kernel, bias) after fusing branches.
        r   rd      )_fuse_bnr   rz   r{   r^   r   r?   rg   rh   r(   r   )rJ   eq_keq_bsmall_ksmall_bs        r!   get_kernel_biasz&ReparamLargeKernelConv.get_kernel_bias  s     ]]4??#7#79K9KL
d4&#}}T__-A-A4??CUCUVGWGODBMM%%4++d.?.??AEFJ D Tzr#   c                    | j                         \  }}t        | j                  | j                  | j                  | j
                  | j                  d      | _        || j                  j                  _	        || j                  j                  _	        | j                  d       t        | d      r| j                  d       yy)a  
        Following works like `RepVGG: Making VGG-style ConvNets Great Again` -
        https://arxiv.org/pdf/2101.03697.pdf. We re-parameterize multi-branched
        architecture used at training time to obtain a plain CNN-like structure
        for inference.
        Tr(   r)   r8   r9   r   r   N)r   r   r&   r'   r(   r)   r8   rB   rY   rZ   r9   r]   r^   )rJ   r   r   s      r!   rb   z%ReparamLargeKernelConv.reparameterize  s     ))+
d)KKLL((;;;;
 )-  %&*#&4&\* 'r#   rz   r{   c                    | j                   }|j                  }|j                  }|j                   }|j                  }|j                  }||z   j                         }||z  j                  dddd      }	||	z  |||z  |z  z
  fS )zMethod to fuse batchnorm layer with conv layer.

        Args:
            conv: Convolutional kernel weights.
            bn: Batchnorm 2d layer.

        Returns:
            Tuple of (kernel, bias) after fusing batchnorm.
        rx   r   )rY   r|   r}   r9   r~   r   r   )
rz   r{   r_   r|   r}   r   r   r~   r   r   s
             r!   r   zReparamLargeKernelConv._fuse_bn  s     nn		wwffS &&(S[!!"aA.z4,"6"<<<<r#   )NFNFNNr1   N)r   r   r   r   r   r   r   r?   r   r>   rf   r   rT   r   r   rb   staticmethodConv2drC   r   r   r   s   @r!   r   r   #  s-    +/ -1#(HKHK HK 	HK
 HK HK #3-HK HK  		*HK !HK 
HKT	 	%,, 	u||U\\'A!B  +. =))== 
u||U\\)	*= =r#   r   FTr&   r'   r0   r+   r.   r1   c                     ||d}t        j                  t        d| |dd|||d|t        d||ddd|||d|t        d||dd|||d|      S )a,  Build convolutional stem with MobileOne blocks.

    Args:
        in_chs: Number of input channels.
        out_chs: Number of output channels.
        inference_mode: Flag to instantiate model in inference mode. Default: ``False``

    Returns:
        nn.Sequential object with stem elements.
    r3      rd   )r&   r'   r(   r)   r0   r+   r.   r   )r&   r'   r(   r)   r   r0   r+   r.   r   )r?   r   r%   )r&   r'   r0   r+   r.   r4   r5   rK   s           r!   convolutional_stemr     s    & U	+B== 		
)-		
 		
 	 
	
)-
	
 
	
 	 		
)-		
 		
-   r#   c                        e Zd ZU dZej
                  j                  e   ed<   	 	 	 	 	 	 dde	de	dede
de
d	df fd
Zdej                  d	ej                  fdZ xZS )	AttentionzMulti-headed Self Attention module.

    Source modified from:
    https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/vision_transformer.py
    
fused_attnNdimhead_dimqkv_bias	attn_drop	proj_dropr1   c                    ||d}t         	|           ||z  dk(  sJ d       || _        ||z  | _        |dz  | _        t               | _        t        j                  ||dz  fd|i|| _	        t        j                  |      | _        t        j                  ||fi || _        t        j                  |      | _        y)a}  Build MHSA module that can handle 3D or 4D input tensors.

        Args:
            dim: Number of embedding dimensions.
            head_dim: Number of hidden dimensions per head. Default: ``32``
            qkv_bias: Use bias or not. Default: ``False``
            attn_drop: Dropout rate for attention tensor.
            proj_drop: Dropout rate for projection tensor.
        r3   r   z#dim should be divisible by head_dimg      r   r9   N)r=   r>   r   	num_headsscaler   r   r?   LinearqkvDropoutr   projr   )
rJ   r   r   r   r   r   r4   r5   rK   rM   s
            r!   r>   zAttention.__init__   s    & /X~"I$II" %
(*99S#'??B?I.IIc3-"-	I.r#   rN   c                 V   |j                   \  }}}}||z  }|j                  d      j                  dd      }| j                  |      j	                  ||d| j
                  | j                        j                  ddddd      }|j                  d      \  }}	}
| j                  rPt        j                  j                  j                  ||	|
| j                  r| j                  j                   nd	      }nL|| j"                  z  }||	j                  dd      z  }|j%                  d
      }| j                  |      }||
z  }|j                  dd      j	                  |||      }| j'                  |      }| j)                  |      }|j                  dd      j	                  ||||      }|S )Nrd   rx   r   r   r   r           )	dropout_p)r   )shapeflatten	transposer   r   r   r   permuteunbindr   rf   r?   rg   scaled_dot_product_attentiontrainingr   pr   softmaxr   r   )rJ   rN   BCHWNr   qkvattns               r!   rT   zAttention.forward   su   WW
1aEIIaL""2r*HHQKWQ1dnndmm<WQ1a# 	
 **Q-1a??##@@1a.2mm$..** A A
 DJJAq{{2r**D<<B<'D>>$'DqAKK1%%aA.IIaLNN1KKB''1a3r#   )    Fr   r   NN)r   r   r   r   rf   jitFinalr   __annotations__r   floatr>   r   rT   r   r   s   @r!   r   r     s    
 		%%
 """// / 	/
 / / 
/@ %,, r#   r   c                        e Zd ZdZej
                  dddddfdededededeej                     d	e	d
e	de	ddf fdZ
dej                  dej                  fdZ xZS )
PatchEmbedz$Convolutional patch embedding layer.FN
patch_sizer)   r&   	embed_dimr0   lkc_use_actr,   r+   r1   c                     |	|
d}t         |           t        j                  t	        d||||dd||r|nd|d	|t        d||ddd||d|      | _        y)	a{  Build patch embedding layer.

        Args:
            patch_size: Patch size for embedding computation.
            stride: Stride for convolutional embedding layer.
            in_chs: Number of channels of input tensor.
            embed_dim: Number of embedding dimensions.
            inference_mode: Flag to instantiate model in inference mode. Default: ``False``
        r3   r   r   N)	r&   r'   r(   r)   r   r   r,   r0   r+   F)r&   r'   r(   r)   r,   r0   r+   r   )r=   r>   r?   r   r   r%   r   )rJ   r   r)   r&   r   r0   r   r,   r+   r4   r5   rK   rM   s               r!   r>   zPatchEmbed.__init__B  s    , /MM" !&'2)-   	 !#-	 	
	r#   rN   c                 (    | j                  |      }|S r   )r   rJ   rN   s     r!   rT   zPatchEmbed.forwards  s    IIaLr#   )r   r   r   r   r?   r   r   r   r   r   r>   rf   r   rT   r   r   s   @r!   r   r   ?  s    . *, % #(/
/
 /
 	/

 /
 BII/
 /
 /
 !/
 
/
b %,, r#   r   c                   <     e Zd Z	 	 	 	 ddededef fdZd Z xZS )LayerScale2dr   init_valuesinplacec           
          t         |           || _        t        j                  |t        j                  |dd||      z        | _        y )Nr   r3   )r=   r>   r   r?   	Parameterrf   onesr   )rJ   r   r   r   r4   r5   rM   s         r!   r>   zLayerScale2d.__init__y  s>     	\\+

31V[`0a"ab
r#   c                 n    | j                   r|j                  | j                        S || j                  z  S r   )r   mul_r   r   s     r!   rT   zLayerScale2d.forward  s(    %)\\qvvdjj!Eq4::~Er#   )h㈵>FNN)	r   r   r   r   r   r   r>   rT   r   r   s   @r!   r   r   x  s<     "&!
c
c 
c 	
cFr#   r   c            	            e Zd ZdZ	 	 	 	 	 ddededee   def fdZde	j                  de	j                  fd	Zdd
Z xZS )RepMixerzReparameterizable token mixer.

    For more details, please refer to our paper:
    `FastViT: A Fast Hybrid Vision Transformer using Structural Reparameterization <https://arxiv.org/pdf/2303.14189.pdf>`_
    r   r(   layer_scale_init_valuer+   c           	         ||d}t         |           || _        || _        || _        |rXt        j                  | j                  | j                  f| j                  d| j                  dz  | j                  dd|| _        yd| _        t        |||fddddd	|| _	        t        |||fddd
|| _
        |t        ||fi || _        yt        j                         | _        y)a  Build RepMixer Module.

        Args:
            dim: Input feature map dimension. :math:`C_{in}` from an expected input of size :math:`(B, C_{in}, H, W)`.
            kernel_size: Kernel size for spatial mixing. Default: 3
            layer_scale_init_value: Initial value for layer scale. Default: 1e-5
            inference_mode: If True, instantiates model in inference mode. Default: ``False``
        r3   r   rd   Tr(   r)   paddingr8   r9   NFr   )r   r-   r.   r/   )r   r-   )r=   r>   r   r(   r+   r?   r   rB   r%   normmixerr   layer_scaler@   )	rJ   r   r(   r   r+   r4   r5   rK   rM   s	           r!   r>   zRepMixer.__init__  s   " /&, "			! !,,((A-xx	! 	!D !%D&	 !&"#	 	DI (  DJ &1#/5K#Rr#R #%;;= r#   rN   r1   c                     | j                   | j                  |      }|S || j                  | j                  |      | j                  |      z
        z   }|S r   )rB   r   r   r   r   s     r!   rT   zRepMixer.forward  sV    (!!!$A  D$$TZZ]TYYq\%ABBAr#   c                 h   | j                   ry| j                  j                          | j                  j                          t	        | j
                  t              r| j                  j                  | j
                  j                  j                  d      | j                  j                  j                  | j                  j                  j                  z
  z  z   }t        j                  | j
                  j                        | j                  j                  j                  | j                  j                  j                  z
  z  }n| j                  j                  | j                  j                  j                  z   | j                  j                  j                  z
  }| j                  j                  j                  | j                  j                  j                  z
  }t        | j                   | j                   | j"                  d| j                   d      | _
        || j                  j                  _        || j                  j                  _        | j'                         D ]  \  }}d|v r|j)                           | j+                  d       | j+                  d       | j+                  d	       y)
ziReparameterize mixer and norm into a single
        convolutional layer for efficient inference.
        Nrx   r   Tr   rB   r   r   r   )r+   r   rb   r   ry   r   r   rv   r   	unsqueezerB   rY   rf   squeezer9   r   r   r(   rZ   r[   r\   r]   )rJ   wbr`   ra   s        r!   rb   zRepMixer.reparameterize  s
    

!!#		  "d&&5

$$t'7'7'='='G'G'K

''..1G1G1N1NN( A d..445

'',,tyy/E/E/J/JJA
 

$$**))001))((//0 
 

'',,tyy/E/E/J/JJA)HHHH((88
 )*  %&'#//1 	JD$%LLN	 	! 'r#   )r   r   FNNr   )r   r   r   r   r   r   r   r   r>   rf   r   rT   rb   r   r   s   @r!   r   r     si      !6:#(9191 91 %-UO	91
 !91v %,, *(r#   r   c                        e Zd ZdZddej
                  dddfdedee   dee   deej                     de
d	df fd
Zdej                  d	dfdZdej                  d	ej                  fdZ xZS )ConvMlpzConvolutional FFN Module.Nr   r&   hidden_channelsr'   r0   dropr1   c                 j   ||d}t         	|           |xs |}|xs |}t        ||fd|dd|| _        t	        j
                  ||fddi|| _         |       | _        t	        j
                  ||fddi|| _        t	        j                  |      | _
        | j                  | j                         y)a_  Build convolutional FFN module.

        Args:
            in_chs: Number of input channels.
            hidden_channels: Number of channels after expansion. Default: None
            out_chs: Number of output channels. Default: None
            act_layer: Activation layer. Default: ``GELU``
            drop: Dropout rate. Default: ``0.0``.
        r3      F)r(   r8   r<   r(   r   N)r=   r>   r   rz   r?   r   fc1rI   fc2r   r   apply_init_weights)
rJ   r&   r   r'   r0   r   r4   r5   rK   rM   s
            r!   r>   zConvMlp.__init__  s    & /#V)3V
 
 
	 99V_J!JrJ;99_gK1KKJJt$	

4%%&r#   mc                     t        |t        j                        rOt        |j                  d       |j
                  +t        j                  j                  |j
                  d       y y y )N{Gz?r   r   )ry   r?   r   r   rY   r9   init	constant_rJ   r  s     r!   r  zConvMlp._init_weights'  sJ    a#!((-vv!!!!&&!, " $r#   rN   c                     | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }|S r   )rz   r   rI   r   r  r   s     r!   rT   zConvMlp.forward-  sU    IIaLHHQKHHQKIIaLHHQKIIaLr#   )r   r   r   r   r?   r   r   r   r   r   r   r>   r  rf   r   rT   r   r   s   @r!   r   r     s    #
 .2%))+#'#' &c]#' c]	#'
 BII#' #' 
#'J-ryy -T - %,, r#   r   c                        e Zd ZdZ	 	 	 	 	 ddedee   deeeeef   f   deddf
 fdZ	d	e
j                  de
j                  fd
ZddZ xZS )RepConditionalPosEnca"  Implementation of conditional positional encoding.

    For more details refer to paper:
    `Conditional Positional Encodings for Vision Transformers <https://arxiv.org/pdf/2102.10882.pdf>`_

    In our implementation, we can reparameterize this module to eliminate a skip connection.
    Nr   dim_outspatial_shaper+   r1   c           
         ||d}t         |           t        |t              rt	        |gdz        }t        |t
              sJ dt        |       d       t        |      dk(  sJ dt        |       d       || _        || _	        |xs || _
        || _        |rQt        j                  | j                  | j                  f| j                  d|d   dz  | j                  dd	|| _        y
d
| _        t        j                  | j                  | j                  |dt        |d   dz        f| j                  dd|| _        y
)at  Build reparameterizable conditional positional encoding

        Args:
            dim: Number of input channels.
            dim_out: Number of embedding dimensions. Default: 768
            spatial_shape: Spatial shape of kernel for positional encoding. Default: (7, 7)
            inference_mode: Flag to instantiate block in inference mode. Default: ``False``
        r3   rd   z/"spatial_shape" must by a sequence or int, get z	 instead.z+Length of "spatial_shape" should be 2, got r   r   Tr   N)r8   r9   )r=   r>   ry   r   tupler   typelenr  r   r  r8   r?   r   rB   pos_enc)	rJ   r   r  r  r+   r4   r5   rK   rM   s	           r!   r>   zRepConditionalPosEnc.__init__@  sj   " /mS)!=/A"56M-/ 	
&'y2	
/ =!Q& 	
}%&i1	
&
 +~# "			! !..%a(A-{{	! 	!D !%D99M!$)*	 {{	 	DLr#   rN   c                 l    | j                   | j                  |      }|S | j                  |      |z   }|S r   )rB   r  r   s     r!   rT   zRepConditionalPosEnc.forward{  s>    (!!!$A  Q!#Ar#   c           
         | j                   | j                  z  }t        j                  | j                   || j                  d   | j                  d   f| j
                  j                  j                  | j
                  j                  j                        }t        | j                         D ].  }d||||z  | j                  d   dz  | j                  d   dz  f<   0 |}|| j
                  j                  z   }| j
                  j                  }t        j                  | j                   | j                  | j                  dt        | j                  d   dz        | j                  d      | _        || j                  j                  _        || j                  j                  _        | j#                         D ]  \  }}d|v r|j%                           | j'                  d       y )	Nr   r   rw   rd   Tr   rB   r  )r   r8   rf   r   r  r  rY   r5   r4   rF   r9   r?   r   r  r   rB   rZ   r[   r\   r]   )	rJ   r   r   r   rv   w_finalb_finalr`   ra   s	            r!   rb   z#RepConditionalPosEnc.reparameterize  s   HH+	{{""1%""1%	 ,,%%++<<&&--	
 txx 	A  I""1%*""1%*,	 !	 dll111,,## IIHHLL****1-23;;
 )0  %&-#//1 	JD$%LLN	 	#r#   )Nr   r   FNNr   )r   r   r   r   r   r   r   r   r   r>   rf   r   rT   rb   r   r   s   @r!   r  r  7  s     &*9?#(99 c]9 !eCHo!56	9
 !9 
9v %,, +$r#   r  c                        e Zd ZdZddej
                  ddddddf	ded	ed
edeej                     dededede
f fdZd Z xZS )RepMixerBlockzImplementation of Metaformer block with RepMixer as token mixer.

    For more details on Metaformer structure, please refer to:
    `MetaFormer Is Actually What You Need for Vision <https://arxiv.org/pdf/2111.11418.pdf>`_
    r         @r   r   FNr   r(   	mlp_ratior0   r   	drop_pathr   r+   c                 R   |	|
d}t         |           t        |f|||d|| _        t	        d|t        ||z        ||d|| _        |t        ||fi || _        nt        j                         | _        |dkD  rt        |      | _        yt        j                         | _        y)a,  Build RepMixer Block.

        Args:
            dim: Number of embedding dimensions.
            kernel_size: Kernel size for repmixer. Default: 3
            mlp_ratio: MLP expansion ratio. Default: 4.0
            act_layer: Activation layer. Default: ``nn.GELU``
            proj_drop: Dropout rate. Default: 0.0
            drop_path: Drop path rate. Default: 0.0
            layer_scale_init_value: Layer scale value at initialization. Default: 1e-5
            inference_mode: Flag to instantiate block in inference mode. Default: ``False``
        r3   )r(   r   r+   r&   r   r0   r   Nr   r   )r=   r>   r   token_mixerr   r   mlpr   r   r?   r@   r   r  )rJ   r   r(   r  r0   r   r  r   r+   r4   r5   rK   rM   s               r!   r>   zRepMixerBlock.__init__  s    2 /#
##9)	

 
  
i0	

 
 "-+C1GN2ND!{{}D09C),R[[]r#   c                     | j                  |      }|| j                  | j                  | j                  |                  z   }|S r   )r!  r  r   r"  r   s     r!   rT   zRepMixerBlock.forward  s=    Qt//<==r#   )r   r   r   r   r?   r   r   r   r   r   r   r>   rT   r   r   s   @r!   r  r    s      !")+"",0#(/S/S /S 	/S
 BII/S /S /S %*/S !/Sbr#   r  c                        e Zd ZdZdej
                  ej                  dddddfdedede	ej                     d	e	ej                     d
ededef fdZd Z xZS )AttentionBlockzImplementation of metaformer block with MHSA as token mixer.

    For more details on Metaformer structure, please refer to:
    `MetaFormer Is Actually What You Need for Vision <https://arxiv.org/pdf/2111.11418.pdf>`_
    r  r   r   Nr   r  r0   
norm_layerr   r  r   c
                    ||	d}
t         |            ||fi |
| _        t        dd|i|
| _        |t        ||fi |
| _        nt        j                         | _        |dkD  rt        |      nt        j                         | _
        t        d|t        ||z        ||d|
| _        |t        ||fi |
| _        nt        j                         | _        |dkD  rt        |      | _        yt        j                         | _        y)a  Build Attention Block.

        Args:
            dim: Number of embedding dimensions.
            mlp_ratio: MLP expansion ratio. Default: 4.0
            act_layer: Activation layer. Default: ``nn.GELU``
            norm_layer: Normalization layer. Default: ``nn.BatchNorm2d``
            proj_drop: Dropout rate. Default: 0.0
            drop_path: Drop path rate. Default: 0.0
            layer_scale_init_value: Layer scale value at initialization. Default: 1e-5
        r3   r   Nr   r   r   )r=   r>   r   r   r!  r   layer_scale_1r?   r@   r   
drop_path1r   r   r"  layer_scale_2
drop_path2)rJ   r   r  r0   r&  r   r  r   r4   r5   rK   rM   s              r!   r>   zAttentionBlock.__init__  s    . /s)b)	$333!-!-c3I!PR!PD!#D1:S(9-bkkm 
i0	

 
 "-!-c3I!PR!PD!#D1:S(9-bkkmr#   c           
          || j                  | j                  | j                  | j                  |                        z   }|| j	                  | j                  | j                  |                  z   }|S r   )r)  r(  r!  r   r+  r*  r"  r   s     r!   rT   zAttentionBlock.forward$  s^     2 243C3CDIIaL3Q RSS 2 2488A; ?@@r#   )r   r   r   r   r?   r   rC   r   r   r   r   r>   rT   r   r   s   @r!   r%  r%    s      #)+*,.."",0-T-T -T BII	-T
 RYY-T -T -T %*-T^r#   r%  c            %           e Zd Zdddddddej                  ej
                  ddd	ddddfd
edededededededede	ej                     dededeej                     deej                     dedeee   ef   de	e   dedef$ fdZd Z xZS )FastVitStageTFr   rd   Nr   r  r   r   r   r  depthtoken_mixer_type
downsamplese_downsampledown_patch_sizedown_stridepos_emb_layerr(   r  r0   r&  proj_drop_ratedrop_path_rater   r   r+   c                 <   t         |           ||d}d| _        |rt        d||||||||d|| _        n ||k(  sJ t        j                         | _        |	 |	|fd|i|| _        nt        j                         | _        g }t        |      D ]r  }|dk(  r'|j                  t        |f|
|||||   ||d|       /|dk(  r&|j                  t        |f||||||   |d	|       Zt        d
j                  |             t        j                  | | _        y)aQ  FastViT stage.

        Args:
            dim: Number of embedding dimensions.
            depth: Number of blocks in stage
            token_mixer_type: Token mixer type.
            kernel_size: Kernel size for repmixer.
            mlp_ratio: MLP expansion ratio.
            act_layer: Activation layer.
            norm_layer: Normalization layer.
            proj_drop_rate: Dropout rate.
            drop_path_rate: Drop path rate.
            layer_scale_init_value: Layer scale value at initialization.
            inference_mode: Flag to instantiate block in inference mode.
        r3   F)r   r)   r&   r   r,   r0   r   r+   Nr+   repmixer)r(   r  r0   r   r  r   r+   	attention)r  r0   r&  r   r  r   z"Token mixer type: {} not supportedr   )r=   r>   grad_checkpointingr   r1  r?   r@   pos_embrF   appendr  r%  
ValueErrorformatr   blocks)rJ   r   r  r/  r0  r1  r2  r3  r4  r5  r(   r  r0   r&  r6  r7  r   r   r+   r4   r5   rK   r@  	block_idxrM   s                           r!   r>   zFastVitStage.__init__+  ss   L 	/"'( 
*"!$#'-
 
DO '>!> kkmDO$(VVSUVDL;;=DLu 	I:-m
 +'',,Y7+A#1
 
 
 "[0n	''),,Y7+A	 	 	 !8??@PQ 3	8 mmV,r#   c                     | j                  |      }| j                  |      }| j                  r6t        j                  j                         st        | j                  |      }|S | j                  |      }|S r   )r1  r<  r;  rf   r   is_scriptingr   r@  r   s     r!   rT   zFastVitStage.forward  s`    OOALLO""599+A+A+Ct{{A.A  AAr#   )r   r   r   r?   r   rC   r   strr   r   r   r   r   r   r   r>   rT   r   r   s   @r!   r.  r.  *  s8     $"'#$ 15 ")+*,..$'8;6: %#(+\-\- \- 	\-
 "\- \-  \- !\- \- $BII.\- \- \- BII\- RYY\- "\-  "$u+u"45!\-" %-UO#\-$ %\-& !'\-|r#   r.  c            3       d    e Zd ZU ej                  j
                  e   ed<   	 ddddddddd	d
dddddddddddej                  ej                  dddfdedeedf   deedf   deedf   deedf   deedf   deedf   dededeeej                      df   ded ed!ed"ed#ed$ed%ed&eded'ed(ed)eej                      d*eej                      d+ed,df2 fd-Zd.ej                   d,dfd/Zej                  j(                  d0        Zej                  j(                  dCd1       Zej                  j(                  dDd2       Zej                  j(                  d,ej                   fd3       ZdEded(ee   fd4Z	 	 	 	 	 dFd5ej4                  d6eeeee   f      d7ed8ed9ed:ed,eeej4                     eej4                  eej4                     f   f   fd;Z	 	 	 dGd6eeee   f   d<ed=efd>Zd5ej4                  d,ej4                  fd?ZdCd5ej4                  d@efdAZ d5ej4                  d,ej4                  fdBZ! xZ"S )Hr   	fork_featr   rd   rd      rd   r9  r9  r9  r9  @            r   r   r   r   )FTTT)FFFF  )NNNNr   rd   r   r   FTg       @avgNin_chanslayers.token_mixers
embed_dims
mlp_ratiosdownsamplesse_downsamplesrepmixer_kernel_sizenum_classespos_embsr3  r4  	drop_rater6  r7  r   r   stem_use_scale_branch	cls_ratioglobal_poolr&  r0   r+   r1   c                    t         (|           ||d}|rdn|	| _        || _        || _        g | _        t        ||d   ||fd|i|| _        |d   }d}t        ||d      }g }t        t        |            D ]  } ||    xs |||    k7  }!t        d$i d|d||    d	||    d
|!d||    d|d|d|
|    d||    d|d||    d|d|d|d||    d|d|d||}"|j                  |"       ||    }|!r|dz  }| xj
                  t        |d|z  d|        gz  c_         t        j                  | | _        t        | j                         | _        |x| _        | _        | j                  rg d| _        t+        | j(                        D ]c  \  }#}$|#dk(  r6t,        j.                  j1                  dd       r	 t        j2                         }%n |||#   fi |}%d|$ }&| j5                  |&|%       e nQt7        |d    |z        x| _        x| _        }'t9        d$|d    |'d!dd|d|dd"	|| _        t=        |'|	f||d#|| _        | jA                  | jB                         y )%Nr3   r   r.   r   T)	stagewiser   r  r/  r1  r2  r3  r4  r5  r0  r(   r  r0   r&  r6  r7  r   r   r+   rd   r   stages.)num_chs	reductionmoduler   r   rd   r   
FORK_LAST3r   rx   r   )	r&   r'   r(   r)   r   r+   r,   r0   r/   )	pool_typer\  r   )"r=   r>   rZ  rF  r_  feature_infor   stemr   rF   r  r.  r=  dictr?   r   stages
num_stagesr:   head_hidden_sizeout_indices	enumerateosenvirongetr@   
add_moduler   r%   
final_convr   headr  r  ))rJ   rR  rS  rT  rU  rV  rW  rX  rY  rZ  r[  r3  r4  r\  r6  r7  r   r   r]  rF  r^  r_  r&  r0   r+   r4   r5   rK   prev_dimr   dprrl  r   r1  stagei_embi_layerlayer
layer_namefinal_featuresrM   s)                                           r!   r>   zFastVit.__init__  sH   : 	/ )1{"& 'qM	

 3
 
	 a='$Os6{# 	eA$QD8z!}+DJ  "1 Qi &	
 -Q/ !0 ( 'qk ".a 1 %Q- $ &  .  #1v  (>!" (#$  .'E* MM% !!}H
$x1u9W^_`^aUb"c!dd7	e8 mmV,dkk*4<<D1 >>  ,D"+D,<,<"= 	3wA:"**..t"D KKME&z%'8?B?E#G9-

E2	3 JMZXZ^^gMgIhhDh 5, !"~&-#"# DO ' &#	
 DI 	

4%%&r#   r  c                    t        |t        j                        rjt        |j                  d       t        |t        j                        r8|j
                  +t        j                  j                  |j
                  d       yyyy)zInit. for classificationr  r  Nr   )ry   r?   r   r   rY   r9   r  r	  r
  s     r!   r  zFastVit._init_weights  sZ    a#!((-!RYY'AFF,>!!!&&!, -?' $r#   c                     t               S r   )setrJ   s    r!   no_weight_decayzFastVit.no_weight_decay  s	    ur#   c                 2    t        d|rd      S g d      S )Nz^stemz^stages\.(\d+)))z^stages\.(\d+).downsampler   )z^stages\.(\d+).pos_embr  )z^stages\.(\d+)\.\w+\.(\d+)N)rj  r@  )rk  )rJ   coarses     r!   group_matcherzFastVit.group_matcher!  s'    (.$
 	
5
 	
r#   c                 4    | j                   D ]	  }||_         y r   )rl  r;  )rJ   enabless      r!   set_grad_checkpointingzFastVit.set_grad_checkpointing,  s     	*A#)A 	*r#   c                 .    | j                   j                  S r   )rv  fcr  s    r!   get_classifierzFastVit.get_classifier1  s    yy||r#   c                 J    || _         | j                  j                  ||       y r   )rZ  rv  reset)rJ   rZ  r_  s      r!   reset_classifierzFastVit.reset_classifier5  s    &		[1r#   rN   indicesr   
stop_early
output_fmtintermediates_onlyc                    |dv sJ d       g }t        t        | j                        |      \  }}	| j                  |      }| j                  dz
  }
t
        j                  j                         s|s| j                  }n| j                  d|	dz    }d}t        |      D ]#  \  }} ||      }||v s|j                  |       % |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   Nr   )r   r  rl  rj  rm  rf   r   rC  rp  r=  ru  )rJ   rN   r  r   r  r  r  intermediatestake_indices	max_indexlast_idxrl  feat_idxry  s                 r!   forward_intermediateszFastVit.forward_intermediates9  s    * Y&D(DD&"6s4;;7G"Qi IIaL??Q&99!!#:[[F[[)a-0F(0 	(OHeaA<'$$Q'	(
   x"A-r#   
prune_norm
prune_headc                     t        t        | j                        |      \  }}| j                  d|dz    | _        |r| j                  dd       |S )z@ Prune layers not required for specified intermediates.
        Nr   r    )r   r  rl  r  )rJ   r  r  r  r  r  s         r!   prune_intermediate_layersz!FastVit.prune_intermediate_layersg  sM     #7s4;;7G"Qikk.9q=1!!!R(r#   c                 <   | j                  |      }g }t        | j                        D ]Q  \  }} ||      }| j                  s|| j                  v s*t        | d|       } ||      }|j                  |       S | j                  r|S | j                  |      }|S )Nr   )rj  rp  rl  rF  ro  getattrr=  ru  )rJ   rN   outsidxblockr&  x_outs          r!   forward_featureszFastVit.forward_featuresu  s    IIaL#DKK0 	'JCaA~~$***!(cU|!<J&qMEKK&	' >>KOOAr#   
pre_logitsc                 N    |r| j                  |d      S | j                  |      S )NT)r  )rv  )rJ   rN   r  s      r!   forward_headzFastVit.forward_head  s$    0:tyyty,L		!Lr#   c                 f    | j                  |      }| j                  r|S | j                  |      }|S r   )r  rF  r  r   s     r!   rT   zFastVit.forward  s3    !!!$>>Ha r#   F)Tr   )NFFr  F)r   FT)#r   r   r   rf   r   r   r   r   r?   rC   r   r   r   rD  r   r   r   r   r>   r  ignorer  r  r  r  r  r   r   r   r  r  r  r  rT   r   r   s   @r!   r   r     s   yyt$$ &2,\*=,4,E/K()#8C#$ "$'$',0 %*.#"$*,..)+#(7z'z' #s(Oz'  S/	z'
 c3hz' eSj)z' tSy)z' "$),z' #&z' z' HRYY/45z' !z' z' z' "z'  "!z'" %*#z'$ %z'& $('z'( )z'* +z', -z'. RYY/z'0 BII1z'2 !3z'8 
9z'x-ryy -T - YY  YY
 
 YY* * YY		  2C 2hsm 2 8<$$',, ||,  eCcN34,  	, 
 ,  ,  !%,  
tELL!5tELL7I)I#JJ	K, ` ./$#	3S	>*  	%,, 5<< "Mell M M %,, r#   c                 2    | dddddt         dt        ddd	|S )
NrP  )r   rM  rM  )   r  g?bicubiczfastvit-license)stem.0.conv_kxk.0.convzstem.0.conv_scale.convzhead.fc)urlrZ  
input_size	pool_sizecrop_pctinterpolationmeanlicenser   
first_conv
classifier)r	   r
   )r  kwargss     r!   _cfgr    s7    #"%$#J  r#   zfastvit_t8.apple_in1kztimm/)	hf_hub_idzfastvit_t12.apple_in1kzfastvit_s12.apple_in1kzfastvit_sa12.apple_in1kzfastvit_sa24.apple_in1kzfastvit_sa36.apple_in1kzfastvit_ma36.apple_in1kgffffff?)r  r  zfastvit_t8.apple_dist_in1kzfastvit_t12.apple_dist_in1kzfastvit_s12.apple_dist_in1kzfastvit_sa12.apple_dist_in1kzfastvit_sa24.apple_dist_in1kzfastvit_sa36.apple_dist_in1kzfastvit_ma36.apple_dist_in1kzfastvit_mci0.apple_mclipzapple/mobileclip_s0_timmzXhttps://docs-assets.developer.apple.com/ml-research/datasets/mobileclip/mobileclip_s0.ptrN  )r   r   r   )      ?r  r  z
apple-amlr)r  r  r  rZ  r  r   r  zfastvit_mci1.apple_mclipzapple/mobileclip_s1_timmzXhttps://docs-assets.developer.apple.com/ml-research/datasets/mobileclip/mobileclip_s1.ptzfastvit_mci2.apple_mclipzapple/mobileclip_s2_timmzXhttps://docs-assets.developer.apple.com/ml-research/datasets/mobileclip/mobileclip_s2.ptr  )r  r  rZ  r  r   r     )r   r   r  )r  r  rZ  r  r   r  r  r  )z!fastvit_mci0.apple_mclip2_dfndr2bz!fastvit_mci2.apple_mclip2_dfndr2bz!fastvit_mci3.apple_mclip2_dfndr2bz!fastvit_mci4.apple_mclip2_dfndr2bc                    d| v r| S d| v rF| j                         D ci c]*  \  }}|j                  d      s|j                  dd      |, c}}S | j                  d|       } d| v rd}nd}d	d
l}d	d
l}g }| j                         D ]C  \  }}|j                  d|      }|s|j                  t        |j                  d                   E t        t        t        |                  }i }	| j                         D ]d  \  }}|r||vr|j                  |d      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd      }|j                  dd |      }|j                  d!      r|j                  d!d"      }|j                  d#d$      }|j                  d%      r|d&k(  rt        |j                   d'      rrt#        |j                   j$                  t&        j(                        rD|j                  d&d(      }|j*                  }t-        j.                  |j0                  d	         |	d)<   n|j                  d%d*      }|j                  d+|      }d,\  }
}|r,t        |j                  d-            }|j3                  ||      }
|
_d.| }d/|
 }|d0z   |v r|j                  |d0z   |d1z         }n5|d2z   |v r|j                  |d2z   |d3z         }n|j                  ||d4z         }||	|<   g |	S c c}}w )5z$ Remap original checkpoints -> timm zstem.0.conv_kxk.0.conv.weightz1module.visual.trunk.stem.0.conv_kxk.0.conv.weightzmodule.visual.trunkzmodule.visual.trunk.r  
state_dictz8image_encoder.model.patch_embed.0.rbr_conv.0.conv.weightzimage_encoder.model.r   Nz^(.*?)network\.(\d+)\.proj.*rd   patch_embedrj  rbr_convrG   	rbr_scalerH   rbr_skiprD   conv_expru  
lkb_originr   convffnr"  z	se.reducezse.fc1z	se.expandzse.fc2zlayer_scale_([0-9])zlayer_scale_\1.gammar   zlayer_scale.gamma	dist_head	head_distzhead.z	head.projr  zhead.fc.weightzhead.fc.biaszhead.fc.z^network\.(\d+))NNr   znetwork.rb  z.projz.downsample.projz.pez.pos_emb.pos_encz.blocks)items
startswithreplacers  rebisectmatchr=  r   grouplistsortedr  subendswithr^   rv  ry   r  r?   r   Trf   r   r   bisect_right)r  modelr   r   prefixr  r  
stage_endsr  out_dict	stage_idxnet_idx
net_prefixstage_prefixs                 r!   checkpoint_filter_fnr    sT   &*4:jHEOEUEUEWTQ[\[g[gh}[~		0"5q8j9JAZO' J  " 318!<c%++a.123 fS_-.JH  " -1Q		&"%A IImV,IIj*-IIk<0IIj*-IIj,/IIlL1IIi'IIk8,IIk8,FF)+BAF::m$		-)<=AIIk;/<< KGEJJ$=*UZZ]]\^\e\eBfIIk+;<CC+0;;qwwqz+B(IIgz2 +Q/'	7%++a.)G++J@I #G9-J$YK0LG#q(IIj72LCU4UVe#q(IIj50,AS2STIIj,*BC[-\ OG @s
   M%M%c                 r    |j                  dd      }t        t        | |ft        t	        d|      d|}|S )Nro  rf  T)flatten_sequentialro  )pretrained_filter_fnfeature_cfg)popr   r   r  rk  )variant
pretrainedr  ro  r  s        r!   _create_fastvitr  M  sJ    **]L9K  2DkJ E Lr#   c           	      N    t        dddd      }t        dd| it        |fi |S )z%Instantiate FastViT-T8 model variant.)rd   rd   r   rd   )0   `        r   r   r   r   rI  rS  rU  rV  rT  r  )
fastvit_t8rk  r  r  r  
model_argss      r!   r  r  Z  s:     %E	J ]J]$zB\U[B\]]r#   c           	      N    t        dddd      }t        dd| it        |fi |S )z&Instantiate FastViT-T12 model variant.rG  rJ  r  rI  r  r  )fastvit_t12r  r  s      r!   r  r  f  :     &E	J ^Z^4
C]V\C]^^r#   c           	      N    t        dddd      }t        dd| it        |fi |S )z&Instantiate FastViT-S12 model variant.rG  rJ  rO  rI  r  r  )fastvit_s12r  r  s      r!   r  r  r  r  r#   c                 v    t        ddddddt        t        d      fd      }t        d
d	| it        |fi |S )z'Instantiate FastViT-SA12 model variant.rG  rJ  rO  Nr  r  r9  r9  r9  r:  rS  rU  rV  r[  rT  r  )fastvit_sa12rk  r   r  r  r  s      r!   r  r  ~  sO     &dG,@PV$WXFJ _j_DD^W]D^__r#   c                 v    t        ddddddt        t        d      fd      }t        d
d	| it        |fi |S )z'Instantiate FastViT-SA24 model variant.)r   r      r   rJ  rO  Nr  r  r  r  r  )fastvit_sa24r   r  s      r!   r  r    O     &dG,@PV$WXFJ _j_DD^W]D^__r#   c                 v    t        ddddddt        t        d      fd      }t        d
d	| it        |fi |S )z'Instantiate FastViT-SA36 model variant.rH  rH     rH  rJ  rO  Nr  r  r  r  r  )fastvit_sa36r   r  s      r!   r  r    r  r#   c                 v    t        ddddddt        t        d      fd      }t        d
d	| it        |fi |S )z'Instantiate FastViT-MA36 model variant.r  )L      i0  i`  rO  Nr  r  r  r  r  )fastvit_ma36r   r  s      r!   r  r    r  r#   c                 z    t        dddddddt        t        d      fdd	
      }t        dd| it        |fi |S )zInstantiate MCi0 model variant.)rd   rH  
   rd   rJ  r  FFTTNr  r  r  TrS  rU  rV  rX  r[  rT  r   r  )fastvit_mci0r   r  s      r!   r  r    sU     &1dG,@PV$WXFJ _j_DD^W]D^__r#   c                 z    t        dddddddt        t        d      fdd	
      }t        dd| it        |fi |S )zInstantiate MCi1 model variant.)r   r     r   rJ  r  r  Nr  r  r  Tr  r  )fastvit_mci1r   r  s      r!   r  r    U     &1dG,@PV$WXFJ _j_DD^W]D^__r#   c                 z    t        dddddddt        t        d      fdd	
      }t        dd| it        |fi |S )zInstantiate MCi2 model variant.)r   r     r   )P      i@  i  r  r  Nr  r  r  Tr  r  )fastvit_mci2r   r  s      r!   r  r    r  r#   c                     t        ddddddddt        t        d      t        t        d      fd	d
t        t        d      d
      }t	        dd| it        |fi |}|S )zInstantiate L model variant.)rd   r  r  r   rd   )r  r  r  r  i   r   r   r   r   r   FFFFFFTTTTNr  r  r9  r9  r9  r:  r:  Tr   r~   F
rS  rU  rV  rX  rW  r[  rT  r   r&  r]  r  )fastvit_mci3rk  r   r  r   r  r  r  r  r  s       r!   r"  r"    s{      ,":3(?(?
 T;D1#!J$ `z`T*E_X^E_`ELr#   c                     t        ddddddddt        t        d      t        t        d      fd	d
t        t        d      d
      }t	        dd| it        |fi |}|S )zInstantiate XL model variant.)rd   r  r  r   r   )rL  rM  rN  i   i   r  r  r  Nr  r  r  Tr   r   Fr!  r  )fastvit_mci4r#  r$  s       r!   r&  r&    s{      .":3(?(?
 T;D1#!J& `z`T*E_X^E_`ELr#   )r  r  )Hrq  	functoolsr   typingr   r   r   r   r   rf   torch.nnr?   	timm.datar	   r
   r   r   timm.layersr   r   r   r   r   r   r   r   r   _builderr   	_featuresr   _manipulater   	_registryr   r   __all__r"   r   r%   r   r   r   r   r   r   r   r   r   r   r   r  r  r%  r.  r   r  default_cfgsr  r  r  r  r  r  r  r  r  r  r  r  r"  r&  r   r#   r!   <module>r2     s   
  5 5   d d
 
 
 + + ' <+&u=RYY u=p[=RYY [=B &(WW$!%444 		?4 	4
 4 ]]4nD		 DN6 6rF299 F"s(ryy s(l5bii 5pv$299 v$r;BII ;|9RYY 9xf299 fR{bii {|" % \&T\& d\& d\& t \& t \& t \& t \&& !$#'\&* "4$+\&0 "4$1\&4 #D%5\&8 #D%9\&< #D%=\&B #D%C\&L ,f|\!M\&Z ,f|\![\&h ,f|\!i\&x *.|* *.|* *.?+* *.?+*g\& \~IX
 ^ ^ _ _ _ _ 	` 	` 	` 	` 	` 	` 	` 	` ` ` ` ` ` `  0  r#   