
    ^jI                     l   d 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 ddl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 G d dej@                        Z! G d dej@                        Z" G d dej@                        Z# G d dej@                        Z$ G d dej@                        Z% G d dej@                        Z&d'dZ' e e'd       e'd       e'd       e'dd       e'dddd       d!      Z(d(d"Z)ed(d#       Z*ed(d$       Z+ed(d%       Z,ed(d&       Z-y))z
InceptionNeXt paper: https://arxiv.org/abs/2303.16900
Original implementation & weights from: https://github.com/sail-sg/inceptionnext
    )partial)ListOptionalTupleUnionTypeNIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)trunc_normal_DropPathcalculate_drop_path_rates	to_2tupleget_paddingSelectAdaptivePool2d   )build_model_with_cfg)feature_take_indices)checkpoint_seq)register_modelgenerate_default_cfgsMetaNeXtc                   L     e Zd ZdZ	 	 	 	 	 	 d	dededededef
 fdZd Z xZS )
InceptionDWConv2dz% Inception depthwise convolution
    in_chssquare_kernel_sizeband_kernel_sizebranch_ratiodilationc                 |   ||d}t         |           t        ||z        }	t        ||      }
t        ||      }t	        j
                  |	|	|f|
||	d|| _        t	        j
                  |	|	d|ffd|fd|f|	d|| _        t	        j
                  |	|	|dff|df|df|	d|| _        |d|	z  z
  |	|	|	f| _	        y )Ndevicedtype)r   )paddingr   groupsr   r      )
super__init__intr   nnConv2d	dwconv_hwdwconv_wdwconv_hsplit_indexes)selfr   r   r   r   r   r"   r#   ddgcsquare_paddingband_padding	__class__s               e/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/inception_next.pyr(   zInceptionDWConv2d.__init__   s    /,&'$%7(K"#3hG&H"XbHDFH 		Q()P%HbPLNP 		%q)P!1%1bPLNP %q2vor2r:    c                     t        j                  || j                  d      \  }}}}t        j                  || j	                  |      | j                  |      | j                  |      fd      S )Nr   )dim)torchsplitr/   catr,   r-   r.   )r0   xx_idx_hwx_wx_hs         r6   forwardzInceptionDWConv2d.forward4   se    ${{1d.@.@aHdCyyNN4 MM#MM#	
 
 	
r7   )r&      g      ?r   NN)	__name__
__module____qualname____doc__r)   floatr(   rB   __classcell__r5   s   @r6   r   r      sV     '($&"';; !$; "	;
  ; ;4
r7   r   c                        e Zd ZdZddej
                  dddddfdedee   dee   deej                     d	eeej                        d
e
def fdZd Z xZS )ConvMlpz MLP using 1x1 convs that keeps spatial dims
    copied from timm: https://github.com/huggingface/pytorch-image-models/blob/v0.6.11/timm/models/layers/mlp.py
    NT        in_featureshidden_featuresout_features	act_layer
norm_layerbiasdropc
                 v   ||	d}
t         |           |xs |}|xs |}t        |      }t        j                  ||fd|d   d|
| _        |r	 ||fi |
nt        j                         | _         |       | _        t        j                  |      | _
        t        j                  ||fd|d   d|
| _        y )Nr!   r   r   )kernel_sizerS   )r'   r(   r   r*   r+   fc1IdentitynormactDropoutrT   fc2)r0   rN   rO   rP   rQ   rR   rS   rT   r"   r#   r1   r5   s              r6   r(   zConvMlp.__init__D   s     /#2{)8[99[/]qtTUw]Z\]9CJ5"5	;JJt$	99_l^PTUVPW^[]^r7   c                     | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }|S N)rW   rY   rZ   rT   r\   r0   r=   s     r6   rB   zConvMlp.forward\   sH    HHQKIIaLHHQKIIaLHHQKr7   )rD   rE   rF   rG   r*   ReLUr)   r   r   ModuleboolrH   r(   rB   rI   rJ   s   @r6   rL   rL   ?   s     .2*.)+48__ &c]_ #3-	_
 BII_ !bii1_ _ _0r7   rL   c                        e Zd ZdZdddej
                   eej                  d      ddd	d	f	d
edede	de
deej                     deej                     de
def fdZddedee	   fdZddefdZ xZS )MlpClassifierHeadz MLP classification head
      avgr&   ư>)epsrM   TNrN   num_classes	pool_type	mlp_ratiorQ   rR   rT   rS   c                    |	|
d}t         |           d| _        || _        t	        ||z        x| _        }|sJ d       t        |d      | _        t        j                  || j                  j                         z  |fd|i|| _         |       | _         ||fi || _        t        j                  ||fd|i|| _        t        j                  |      | _        y )Nr!   FCannot disable poolingTrj   flattenrS   )r'   r(   use_convrN   r)   num_featuresr   global_poolr*   Linear	feat_multrW   rZ   rY   r\   r[   rT   )r0   rN   ri   rj   rk   rQ   rR   rT   rS   r"   r#   r1   rO   r5   s                r6   r(   zMlpClassifierHead.__init__i   s     /&.1)k2I.JJO222y/)TR99[4+;+;+E+E+GGj_cjgij;5"5	99_kKKKJJt$	r7   c                     ||sJ d       t        |d      | _        |dkD  r&t        j                  | j                  |      | _        y t        j
                         | _        y )Nrm   Trn   r   )r   rr   r*   rs   rq   rX   r\   )r0   ri   rj   s      r6   resetzMlpClassifierHead.reset   sT     66693iQUVD@Ka299T..<UWU`U`Ubr7   
pre_logitsc                     | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }|r|S | j                  |      S r^   )rr   rW   rZ   rY   rT   r\   r0   r=   rw   s      r6   rB   zMlpClassifierHead.forward   s[    QHHQKHHQKIIaLIIaLq/DHHQK/r7   r^   F)rD   rE   rF   rG   r*   GELUr   	LayerNormr)   strrH   r   ra   rb   r(   r   rv   rB   rI   rJ   s   @r6   rd   rd   e   s      $" )+*1",,D*I%% % 	%
 % BII% RYY% % %8c c# c0T 0r7   rd   c                        e Zd ZdZdeej                  edej                  ddddf
de	de	d	e
ej                     d
e
ej                     de
ej                     dede
ej                     dedef fdZd Z xZS )MetaNeXtBlockz MetaNeXtBlock Block
    Args:
        dim (int): Number of input channels.
        drop_path (float): Stochastic depth rate. Default: 0.0
        ls_init_value (float): Init value for Layer Scale. Default: 1e-6.
    r      rg   rM   Nr9   r   token_mixerrR   	mlp_layerrk   rQ   ls_init_value	drop_pathc                 j   |
|d}t         |            ||fd|i|| _         ||fi || _         ||t	        ||z        fd|i|| _        |r,t        j                  |t        j                  |fi |z        nd | _
        |	dkD  rt        |	      | _        y t        j                         | _        y )Nr!   r   rQ   rM   )r'   r(   r   rY   r)   mlpr*   	Parameterr:   onesgammar   rX   r   )r0   r9   r   r   rR   r   rk   rQ   r   r   r"   r#   r1   r5   s                r6   r(   zMetaNeXtBlock.__init__   s     /&sDXDDs)b)	S#i#o"6R)RrRLYR\\-%**S2GB2G"GH_c
09B),BKKMr7   c                 
   |}| j                  |      }| j                  |      }| j                  |      }| j                  -|j	                  | j                  j                  dddd            }| j                  |      |z   }|S )Nr   )r   rY   r   r   mulreshaper   )r0   r=   shortcuts      r6   rB   zMetaNeXtBlock.forward   sv    QIIaLHHQK::!djj((B156ANN1(r7   )rD   rE   rF   rG   r   r*   BatchNorm2drL   r{   r)   r   ra   rH   r(   rB   rI   rJ   s   @r6   r   r      s     +<*,..)0 )+#'!RR R bii	R
 RYYR BIIR R BIIR !R R,r7   r   c                        e Zd Zdddddeej
                  ddddfdededed	ed
eeef   dee	e
      de
deej                     deej                     deeej                        de
f fdZd Z xZS )MetaNeXtStage   )r   r   N      ?r   r   out_chsstridedepthr   drop_path_ratesr   r   rQ   rR   rk   c                    ||d}t         |           d| _        |dkD  s|d   |d   k7  r?t        j                   |
|fi |t        j
                  ||fd||d   d|      | _        nt        j                         | _        |xs dg|z  }g }t        |      D ]+  }|j                  t        d	||d   ||   |||	|
|d|       - t        j                  | | _        y )
Nr!   Fr   r   r   )rV   r   r   rM   )r9   r   r   r   r   rQ   rR   rk    )r'   r(   grad_checkpointingr*   
Sequentialr+   
downsamplerX   rangeappendr   blocks)r0   r   r   r   r   r   r   r   r   rQ   rR   rk   r"   r#   r1   stage_blocksir5   s                    r6   r(   zMetaNeXtStage.__init__   s     /"'A:!3 mm6(R(		 !"!%a[ 
DO !kkmDO)9bTE\u 	A 
!!!)!,+'#%#
! 
! 
	 mm\2r7   c                     | j                  |      }| j                  r6t        j                  j	                         st        | j                  |      }|S | j                  |      }|S r^   )r   r   r:   jitis_scriptingr   r   r_   s     r6   rB   zMetaNeXtStage.forward   sS    OOA""599+A+A+Ct{{A.A  AAr7   )rD   rE   rF   r   r*   r{   r)   r   r   r   rH   r   ra   r(   rB   rI   rJ   s   @r6   r   r      s    
 (.59#&+<)+48 0303 03 	03
 03 CHo03 &d5k203 !03 bii03 BII03 !bii103 03dr7   r   c                   X    e Zd ZdZddddddeej                  ej                  dd	d	d
ddfdedede	dede
edf   de
edf   deeej                     eeej                        f   deej                     deej                     deee
edf   f   dededef fdZd Zej&                  j(                  d0d       Zej&                  j(                  dej                  fd       Zd1dedee	   fdZej&                  j(                  d2d        Zej&                  j(                  d!        Z	 	 	 	 	 d3d"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	 	 	 d4d#eeee   f   d)ed*efd+Zd, Zd0d-efd.Z d/ Z! xZ"S )5r   a   MetaNeXt
        A PyTorch impl of : `InceptionNeXt: When Inception Meets ConvNeXt` - https://arxiv.org/abs/2303.16900

    Args:
        in_chans (int): Number of input image channels. Default: 3
        num_classes (int): Number of classes for classification head. Default: 1000
        depths (tuple(int)): Number of blocks at each stage. Default: (3, 3, 9, 3)
        dims (tuple(int)): Feature dimension at each stage. Default: (96, 192, 384, 768)
        token_mixers: Token mixer function. Default: nn.Identity
        norm_layer: Normalization layer. Default: nn.BatchNorm2d
        act_layer: Activation function for MLP. Default: nn.GELU
        mlp_ratios (int or tuple(int)): MLP ratios. Default: (4, 4, 4, 3)
        drop_rate (float): Head dropout rate
        drop_path_rate (float): Stochastic depth rate. Default: 0.
        ls_init_value (float): Init value for Layer Scale. Default: 1e-6.
    r&   re   rf       r&   r&   	   r&   `        i   )r   r   r   r&   rM   rg   Nin_chansri   rr   output_stridedepths.dimstoken_mixersrR   rQ   
mlp_ratios	drop_ratedrop_path_rater   c                    t         |           ||d}t        |      }t        |t        t
        f      s|g|z  }t        |
t        t
        f      s|
g|z  }
|| _        || _        || _        || _	        g | _
        t        j                  t        j                  ||d   fddd| ||d   fi |      | _        t        ||d      }|d   }d}d}t        j                         | _        t#        |      D ]  }|dk(  s|dkD  rdnd}||k\  r|dkD  r||z  }d}||z  }|d	v rdnd}||   }| j                   j%                  t'        ||f|dkD  r|nd||f||   ||   ||	||   ||
|   d
	|       |}| xj                  t)        ||d|       gz  c_
         || _        t-        | j*                  |f| j                  |d|| _        | j.                  j*                  | _        | j3                  | j4                         y )Nr!   r   r   )rV   r   T)	stagewiser   r   )r   r   )	r   r   r   r   r   rQ   r   rR   rk   zstages.)num_chs	reductionmodule)rj   rT   )r'   r(   len
isinstancelisttupleri   r   rr   r   feature_infor*   r   r+   stemr   stagesr   r   r   dictrq   rd   headhead_hidden_sizeapply_init_weights)r0   r   ri   rr   r   r   r   r   rR   rQ   r   r   r   r   r"   r#   r1   	num_stagedp_ratesprev_chscurr_strider   r   r   first_dilationr   r5   s                             r6   r(   zMetaNeXt.__init__  sA   $ 	/K	,u6(>I5L*tUm4$	1J& &"MMIIhQGQqGBGtAw%"%
	
 -^VtT7mmoy! 	gA%*a!eQFm+
F"6!K"*f"4Q!N1gGKK}  "#QvA((3Qi (+#(O%$Q-     H$x;Y`ab`cWd"e!ff/	g0 %%d&7&7wPTP`P`gpwtvw	 $		 6 6

4%%&r7   c                     t        |t        j                  t        j                  f      rOt	        |j
                  d       |j                  +t        j                  j                  |j                  d       y y y )Ng{Gz?)stdr   )	r   r*   r+   rs   r   weightrS   init	constant_)r0   ms     r6   r   zMetaNeXt._init_weightsS  sS    a"))RYY/0!((,vv!!!!&&!, " 1r7   c                 2    t        d|rd      S ddg      S )Nz^stemz^stages\.(\d+))z^stages\.(\d+)\.downsample)r   )z^stages\.(\d+)\.blocks\.(\d+)N)r   r   )r   )r0   coarses     r6   group_matcherzMetaNeXt.group_matcherY  s/    (.$
 	
 685
 	
r7   returnc                 .    | j                   j                  S r^   )r   r\   r0   s    r6   get_classifierzMetaNeXt.get_classifierc  s    yy}}r7   c                 J    || _         | j                  j                  ||       y r^   )ri   r   rv   )r0   ri   rr   s      r6   reset_classifierzMetaNeXt.reset_classifierg  s    &		[1r7   c                 4    | j                   D ]	  }||_         y r^   )r   r   )r0   enabless      r6   set_grad_checkpointingzMetaNeXt.set_grad_checkpointingk  s     	*A#)A 	*r7   c                     t               S r^   )setr   s    r6   no_weight_decayzMetaNeXt.no_weight_decayp  s	    ur7   r=   indicesrY   
stop_early
output_fmtintermediates_onlyc                 r   |dv sJ d       g }t        t        | j                        |      \  }}	| j                  |      }t        j
                  j                         s|s| j                  }
n| j                  d|	dz    }
t        |
      D ]#  \  }} ||      }||v s|j                  |       % |r|S ||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.Nr   )	r   r   r   r   r:   r   r   	enumerater   )r0   r=   r   rY   r   r   r   intermediatestake_indices	max_indexr   feat_idxstages                r6   forward_intermediateszMetaNeXt.forward_intermediatest  s    * Y&D(DD&"6s4;;7G"Qi IIaL99!!#:[[F[[)a-0F(0 	(OHeaA<'$$Q'	(
   -r7   
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   rf   )r   r   r   r   )r0   r   r   r   r   r   s         r6   prune_intermediate_layersz"MetaNeXt.prune_intermediate_layers  sM     #7s4;;7G"Qikk.9q=1!!!U+r7   c                 J    | j                  |      }| j                  |      }|S r^   )r   r   r_   s     r6   forward_featureszMetaNeXt.forward_features  s!    IIaLKKNr7   rw   c                 N    |r| j                  ||      S | j                  |      S )N)rw   )r   ry   s      r6   forward_headzMetaNeXt.forward_head  s%    6@tyyzy2RdiiPQlRr7   c                 J    | j                  |      }| j                  |      }|S r^   )r   r   r_   s     r6   rB   zMetaNeXt.forward  s'    !!!$a r7   rz   r^   )T)NFFr   F)r   FT)#rD   rE   rF   rG   r   r*   r   r{   r)   r}   r   r   r   ra   r   rH   r(   r   r:   r   ignorer   r   r   r   r   r   Tensorrb   r   r   r   r   rB   rI   rJ   s   @r6   r   r      s   & #$!#&2$7J[*,..)+6B!$&#'!E'E' E' 	E'
 E' #s(OE' S/E'  RYYd299o1F FGE' RYYE' BIIE' c5c?23E' E' "E' !E'N- YY
 
 YY		  2C 2hsm 2 YY* * YY  8<$$',( ||(  eCcN34(  	( 
 (  (  !%(  
tELL!5tELL7I)I#JJ	K( X ./$#	3S	>*  	
S$ Sr7   c                 2    | dddddt         t        dddd	|S )
Nre   )r&      r   )   r   g      ?bicubiczstem.0zhead.fc2z
apache-2.0)urlri   
input_size	pool_sizecrop_pctinterpolationmeanr   
first_conv
classifierlicenser	   )r   kwargss     r6   _cfgr    s3    =vI%.Bj  r7   ztimm/)	hf_hub_idgffffff?)r  r   )r&   r   r   )   r	  r   )r  r   r   r   )zinception_next_atto.sail_in1kzinception_next_tiny.sail_in1kzinception_next_small.sail_in1kzinception_next_base.sail_in1kz!inception_next_base.sail_in1k_384c                 D    t        t        | |fdt        dd      i|}|S )Nfeature_cfg)r   r   r   r&   T)out_indicesflatten_sequential)r   r   r   )variant
pretrainedr  models       r6   _create_inception_nextr    s3     ':\dK E
 Lr7   c           	      l    t        ddt        t        dd            }t        dd| it        |fi |S )	N)r   r      r   )(   P      i@  r   g      ?)r   r   r   r   r   r  )inception_next_atto)r   r   r   r  r  r  
model_argss      r6   r  r    sB    "4.QUVJ "mJmRVWaRlekRlmmr7   c           	      T    t        ddt              }t        dd| it        |fi |S )Nr   r   r  r  )inception_next_tinyr   r   r  r  s      r6   r  r    s7    "5&J "mJmRVWaRlekRlmmr7   c           	      T    t        ddt              }t        dd| it        |fi |S )Nr&   r&      r&   r   r  r  )inception_next_smallr  r  s      r6   r!  r!    s7    #6&J "nZnSWXbSmflSmnnr7   c           	      T    t        ddt              }t        dd| it        |fi |S )Nr  )      i   i   r  r  )inception_next_baser  r  s      r6   r%  r%    s7    #8&J "mJmRVWaRlekRlmmr7   ) rz   ).rG   	functoolsr   typingr   r   r   r   r   r:   torch.nnr*   	timm.datar
   r   timm.layersr   r   r   r   r   r   _builderr   	_featuresr   _manipulater   	_registryr   r   __all__ra   r   rL   rd   r   r   r   r  default_cfgsr  r  r  r!  r%  r   r7   r6   <module>r2     s[  
  5 5   A x x * + ' <,&
		 &
R#bii #L-0		 -0`&BII &R9BII 9x}ryy }@	 %%)& &*& '+' &*&
 *. Hs*%& 4 n n n n o o n nr7   