
    ^j                     V   d Z ddlmZ ddlmZm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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dgZ(e G d d             Z) G d dejT                        Z+dde,de-de
fdZ. G d dejT                        Z/e" G d dejT                               Z0	 	 	 	 	 	 dde1de1de,d ee
   d!ee
   d"e2deejf                  e1ee,e	f   f   fd#Z4 e5dd$d%d&d'd(d)d*d+d,d-d.d/d0d12      Z6 G d3 dejT                        Z7	 	 	 	 	 dd4ee1d5f   d6ee1d5f   d7ee1   d!e,d8ee,   d9eee,e	f      de)fd:Z8dd4ee1d5f   d6ee1d5f   de)fd;Z9	 	 	 	 	 	 	 dd4ee1d5f   d6ee1d5f   d7e1d=e-d>e-d!e,d8e,d9eee,e	f      de)fd?Z:	 	 	 d	d4ee1d5f   d6ee1d5f   d!e,d@e2de)f
dAZ; e5d
i dB e;dCD      dE e;dFD      dG e;dHD      dI e;dJD      dK e;dLD      dM e;dND      dO e;dPD      dQ e:dCD      dR e:dFD      dS e:dHD      dT e:dJD      dU e:dLD      dV e:dND      dW e:dPD      dX e:dYD      dZ e:dCd[d\d] e5d]d^_      d`a      db e:dCd[d\d]dc e5       d`d      de e:dFdfd\d]dc e5       d`d      dg e:dHdfd\d]dc e5       d`d      dh e:dJdfd\d]dc e5       d`d      di e9djD      dk e9dlD      dm e9dndop      dq e9drdsp      dt e9dudvp      dw e9dxdyp      dz e8d{D      d| e8d}D      d~ e8dD      d e8d{d< e5d            d e8d}d< e5d            d e8dd< e5d            d e8d{dc e5             d e8d}dc e5             d e8ddc e5             d e:ddd[d^d] e5d]d^_      d`      Z<dde,de2de	de7fdZ=dde,de	dee,e	f   fdZ> e&i d e>ddddddd      d e>ddddddd      d e>ddddddd      d e>ddddddd      d e>ddddddd      d e>ddddddd      d e>ddddddd      dQ e>dddd      dR e>dddd      dS e>dddd      dT e>dddd      dU e>dddd      dV e>dddd      dW e>dddd      dX e>dddd      d e>ddddddī      d e>ddddddī      i d e>ddddddī      d e>ddddddī      dh e>ddddd̫      di e>dddddͬΫ      d e>ddddddͬѫ      dm e>dddddͬΫ      dq e>dddddͬΫ      dt e>dddddͬΫ      dw e>dddddͬΫ      dz e>ddͬ֫      d e>dddddddͬ٫      d~ e>ddͬ֫      d e>ddͬ֫      d e>ddͬ֫      d e>ddͬ֫      d e>ddͬ֫      d e>ddͬ֫       e>ddͬ֫       e>ddddddݬޫ      dߜ      Z?e'dde2de	de7fd       Z@e'dde2de	de7fd       ZAe'dde2de	de7fd       ZBe'dde2de	de7fd       ZCe'dde2de	de7fd       ZDe'dde2de	de7fd       ZEe'dde2de	de7fd       ZFe'dde2de	de7fd       ZGe'dde2de	de7fd       ZHe'dde2de	de7fd       ZIe'dde2de	de7fd       ZJe'dde2de	de7fd       ZKe'dde2de	de7fd       ZLe'dde2de	de7fd       ZMe'dde2de	de7fd       ZNe'dde2de	de7fd       ZOe'dde2de	de7fd       ZPe'dde2de	de7fd       ZQe'dde2de	de7fd       ZRe'dde2de	de7fd       ZSe'dde2de	de7fd       ZTe'dde2de	de7fd       ZUe'dde2de	de7fd       ZVe'dde2de	de7fd       ZWe'dde2de	de7fd       ZXe'dde2de	de7fd       ZYe'dde2de	de7fd       ZZe'dde2de	de7fd       Z[e'dde2de	de7fd       Z\e'dde2de	de7fd       Z]e'dde2de	de7fd       Z^e'dde2de	de7fd       Z_e'dde2de	de7fd        Z`e'dde2de	de7fd       Zae'dde2de	de7fd       Zbe'dde2de	de7fd       Zcy(  a   Normalization Free Nets. NFNet, NF-RegNet, NF-ResNet (pre-activation) Models

Paper: `Characterizing signal propagation to close the performance gap in unnormalized ResNets`
    - https://arxiv.org/abs/2101.08692

Paper: `High-Performance Large-Scale Image Recognition Without Normalization`
    - https://arxiv.org/abs/2102.06171

Official Deepmind JAX code: https://github.com/deepmind/deepmind-research/tree/master/nfnets

Status:
* These models are a work in progress, experiments ongoing.
* Pretrained weights for two models so far, more to come.
* Model details updated to closer match official JAX code now that it's released
* NF-ResNet, NF-RegNet-B, and NFNet-F models supported

Hacked together by / copyright Ross Wightman, 2021.
    )OrderedDict)	dataclassreplace)partial)AnyCallableDictOptionalTupleNIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)
ClassifierHeadDropPathcalculate_drop_path_ratesAvgPool2dSameScaledStdConv2dScaledStdConv2dSameget_act_layer
get_act_fnget_attnmake_divisible   )build_model_with_cfg)register_notrace_module)checkpoint_seq)generate_default_cfgsregister_modelNormFreeNetNfCfgc                   n   e Zd ZU dZeeeeef   ed<   e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   ed
<   dZee
   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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!<   y)"r    z.Configuration for Normalization-Free Networks.depthschannelsg?alpha3x3	stem_typeNstem_chs
group_size
attn_layerattn_kwargs       @	attn_gain      ?width_factor      ?bottle_ratior   num_features   ch_divFreg
extra_convgamma_in_actsame_paddinggh㈵>std_conv_epsskipinitzero_init_fcsilu	act_layer) __name__
__module____qualname____doc__r   int__annotations__r$   floatr&   strr'   r
   r(   r)   r*   r	   r   r,   r.   r0   r1   r3   r4   boolr5   r6   r7   r8   r9   r:   r<        \/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/nfnet.pyr    r    &   s   8#sC$%%Cc3&''E5Is"Hhsm" $J$ $J$,0K$sCx.)0IuL%L%L#FCOCJL$L$L%HdL$IsrG   c                   j     e Zd ZdZd	dededef fdZdej                  dej                  fdZ
 xZS )
GammaActz.Activation function with gamma scaling factor.act_typegammainplacec                 ^    t         |           t        |      | _        || _        || _        y)zInitialize GammaAct.

        Args:
            act_type: Type of activation function.
            gamma: Scaling factor for activation output.
            inplace: Whether to perform activation in-place.
        N)super__init__r   act_fnrL   rM   )selfrK   rL   rM   	__class__s       rH   rP   zGammaAct.__init__C   s*     	 *
rG   xreturnc                 n    | j                  || j                        j                  | j                        S )zzForward pass.

        Args:
            x: Input tensor.

        Returns:
            Scaled activation output.
        rM   )rQ   rM   mul_rL   rR   rT   s     rH   forwardzGammaAct.forwardP   s*     {{1dll{388DDrG   )relur-   F)r=   r>   r?   r@   rD   rC   rE   rP   torchTensorrZ   __classcell__rS   s   @rH   rJ   rJ   @   s>    8 e D 	E 	E%,, 	ErG   rJ   r-   rK   rL   rU   c                 2     ddt         dt        f fd}|S )zCreate activation function factory with gamma scaling.

    Args:
        act_type: Type of activation function.
        gamma: Scaling factor for activation output.

    Returns:
        Activation function factory.
    rM   rU   c                      t        |       S )N)rL   rM   )rJ   )rM   rK   rL   s    rH   _createzact_with_gamma.<locals>._createf   s    w??rG   F)rE   rJ   )rK   rL   rb   s   `` rH   act_with_gammard   \   s    @ @( @NrG   c                        e Zd ZdZdddeddfdededededee   d	ef fd
Zde	j                  de	j                  fdZ xZS )DownsampleAvgzEAvgPool downsampling as in 'D' ResNet variants with dilation support.r   Nin_chsout_chsstridedilationfirst_dilation
conv_layerc	                    t         |           |dk(  r|nd}	|dkD  s|dkD  r2|	dk(  r|dkD  rt        nt        j                  }
 |
d|	dd      | _        nt        j                         | _         |||dd||      | _        y)a7  Initialize DownsampleAvg.

        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            stride: Stride for downsampling.
            dilation: Dilation rate.
            first_dilation: First dilation rate (unused).
            conv_layer: Convolution layer type.
        r      TF)	ceil_modecount_include_pad)ri   devicedtypeN)rO   rP   r   nn	AvgPool2dpoolIdentityconv)rR   rg   rh   ri   rj   rk   rl   rq   rr   
avg_strideavg_pool_fnrS   s              rH   rP   zDownsampleAvg.__init__n   s{    * 	'1}V!
A:A+5?x!|-QSQ]Q]K#AzTUZ[DIDIvw!FRWX	rG   rT   rU   c                 B    | j                  | j                  |            S )ztForward pass.

        Args:
            x: Input tensor.

        Returns:
            Downsampled tensor.
        )rw   ru   rY   s     rH   rZ   zDownsampleAvg.forward   s     yy1&&rG   )r=   r>   r?   r@   r   rA   r
   r   rP   r\   r]   rZ   r^   r_   s   @rH   rf   rf   k   s    O ,0#2YY Y 	Y
 Y %SMY !Y<	' 	'%,, 	'rG   rf   c            %            e Zd ZdZddddddddddddddded	ddfd
edee   dededee   dedededee   dededededee	   dedee	   de	def$ fdZ
dej                  dej                  fdZ xZS ) NormFreeBlockz-Normalization-Free pre-activation block.
    Nr   r-         ?TFr+           rg   rh   ri   rj   rk   r$   betar0   r(   r3   r4   r5   r9   r)   r,   r<   rl   drop_path_ratec                 t   ||d}t         |           |xs |}|xs |}t        |r||z  n||z  |
      }|	sdn||	z  }|	r|	|
z  dk(  r|	|z  }|| _        || _        || _        ||k7  s
|dk7  s||k7  rt        ||f||||d|| _        nd| _         |       | _         |||dfi || _	         |d      | _
         |||df|||d	|| _        |r# |d      | _         |||dfd||d	|| _        nd| _        d| _        |r| ||fi || _        nd| _         |       | _         |||dfd
|rdndi|| _        |s| ||fi || _        nd| _        |dkD  rt%        |      nt'        j(                         | _        |r.t'        j,                  t/        j0                  di |      | _        yd| _        y)a  Initialize NormFreeBlock.

        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            stride: Stride for convolution.
            dilation: Dilation rate.
            first_dilation: First dilation rate.
            alpha: Alpha scaling factor for residual.
            beta: Beta scaling factor for pre-activation.
            bottle_ratio: Bottleneck ratio.
            group_size: Group convolution size.
            ch_div: Channel divisor for rounding.
            reg: Use RegNet-style configuration.
            extra_conv: Add extra 3x3 convolution.
            skipinit: Use skipinit initialization.
            attn_layer: Attention layer type.
            attn_gain: Attention gain factor.
            act_layer: Activation layer type.
            conv_layer: Convolution layer type.
            drop_path_rate: Stochastic depth drop rate.
        rq   rr   r   r   )ri   rj   rk   rl   NTrW      )ri   rj   groups	gain_initr-   r~   )r~   )rO   rP   r   r$   r   r,   rf   
downsampleact1conv1act2conv2act2bconv2battnact3conv3	attn_lastr   rs   rv   	drop_path	Parameterr\   tensorskipinit_gain)rR   rg   rh   ri   rj   rk   r$   r   r0   r(   r3   r4   r5   r9   r)   r,   r<   rl   r   rq   rr   ddmid_chsr   rS   s                           rH   rP   zNormFreeBlock.__init__   s   Z /'38#V #,!67\CY[ab$'Z*?*v-2 6)G
	"W!x>/I+ !-% DO #DOK	9b9
d+	!qF^djqnpq
"40DJ$WgqkX^dkhjkDKDJDK:)"71b1DIDIK	!\XrSU\Y[\
z-'626DN!DN5Ca5G.1R[[]EMR\\%,,*@R*@ASWrG   rT   rU   c                    | j                  |      | j                  z  }|}| j                  | j                  |      }| j                  |      }| j	                  | j                  |            }| j                   | j                  | j                  |            }| j                  | j                  | j                  |      z  }| j                  | j                  |            }| j                  | j                  | j                  |      z  }| j                  |      }| j                  |j                  | j                         || j                   z  |z   }|S )zoForward pass.

        Args:
            x: Input tensor.

        Returns:
            Output tensor.
        )r   r   r   r   r   r   r   r   r   r,   r   r   r   r   r   rX   r$   )rR   rT   outshortcuts       rH   rZ   zNormFreeBlock.forward   s    iilTYY& ??&s+H jjojj3(;;"++djjo.C99 ..499S>1Cjj3(>>%..4>>##66CnnS!)HHT''(DJJ)
rG   )r=   r>   r?   r@   r   rA   r
   rC   rE   r   rP   r\   r]   rZ   r^   r_   s   @rH   r|   r|      sR    &*,0"&(,$"-1",0#2$&+\X\X c]\X 	\X
 \X %SM\X \X \X  \X !\X \X \X \X \X !*\X  !\X"  )#\X$ !%\X& "'\X| %,, rG   r|    rg   rh   r&   rl   r<   preact_featurec                    ||d}d}	t        |dd      }
t               }|dv sJ d|v rd|v r+d|vsJ |d	z  |d
z  |dz  |f}d}d
}	t        |dz  dd      }
n1d|v rd|z  d	z  |dz  |f}n|dz  |dz  |f}d}t        |dz  dd      }
t        |      dz
  }t        t	        ||            D ]7  \  }\  }} || |fd|d||d|dz    <   ||k7  r |d      |d|dz    <   |} 9 n%d|v r || |fddd||d<   n || |fddd||d<   d|v rt        j                  ddd      |d<   d
}	t        j                  |      |	|
fS )a  Create stem module for NFNet models.

    Args:
        in_chs: Input channels.
        out_chs: Output channels.
        stem_type: Type of stem ('', 'deep', 'deep_tiered', 'deep_quad', '3x3', '7x7', etc.).
        conv_layer: Convolution layer type.
        act_layer: Activation layer type.
        preact_feature: Use pre-activation feature.

    Returns:
        Tuple of (stem_module, stem_stride, stem_feature_info).
    r   rn   	stem.convnum_chs	reductionmodule)	r   deepdeep_tiered	deep_quadr%   7x7	deep_pool3x3_pool7x7_poolr   quadru   r2      )rn   r   r   rn   z
stem.conv3tieredr   )rn   r   r   z
stem.conv2r   )kernel_sizeri   rw   TrW   actr%      )ri   padding)dictr   len	enumerateziprs   	MaxPool2d
Sequential)rg   rh   r&   rl   r<   r   rq   rr   r   stem_stridestem_featurestemr'   strideslast_idxicss                     rH   create_stemr     s   . U	+BK1[IL=DssssY***1glGqL'JH"GK1,WL9$K1,glGD#qL'Q,@G1,WLx=1$"3x#9: 	IAv1#-fa#WQq#WTV#WD4Aw H}&/&=s1q5']#F		
 
)	!&'QqQbQV "&'QqQbQV||Aa;V==\99rG   g   `U?g   yX?g   \9?g   `aK?g   ?g    ?g    `l?g   `i?g   |?g    7@g   -?g   @g   `?g   ?)identityceluelugelu
leaky_relulog_sigmoidlog_softmaxr[   relu6selusigmoidr;   softsignsoftplustanhc                       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	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e
j,                  de
j,                  fdZdde
j,                  dede
j,                  fdZde
j,                  de
j,                  fdZ xZS )r   a*   Normalization-Free Network

    As described in :
    `Characterizing signal propagation to close the performance gap in unnormalized ResNets`
        - https://arxiv.org/abs/2101.08692
    and
    `High-Performance Large-Scale Image Recognition Without Normalization` - https://arxiv.org/abs/2102.06171

    This model aims to cover both the NFRegNet-Bx models as detailed in the paper's code snippets and
    the (preact) ResNet models described earlier in the paper.

    There are a few differences:
        * channels are rounded to be divisible by 8 by default (keep tensor core kernels happy),
            this changes channel dim and param counts slightly from the paper models
        * activation correcting gamma constants are moved into the ScaledStdConv as it has less performance
            impact in PyTorch when done with the weight scaling there. This likely wasn't a concern in the JAX impl.
        * a config option `gamma_in_act` can be enabled to not apply gamma in StdConv as described above, but
            apply it in each activation. This is slightly slower, numerically different, but matches official impl.
        * skipinit is disabled by default, it seems to have a rather drastic impact on GPU memory use and throughput
            for what it is/does. Approx 8-10% throughput loss.
    Ncfgnum_classesin_chansglobal_pooloutput_stride	drop_rater   kwargsc
                    t         "|           ||	d}|| _        || _        || _        d| _        t        |fi |
}|j                  t        v sJ d|j                   d       |j                  rt        nt        }|j                  r@t        |j                  t        |j                           }t        ||j                        }n>t!        |j                        }t        |t        |j                     |j                        }|j"                  r)t        t%        |j"                        fi |j&                  nd}t)        |j*                  xs |j,                  d	   |j.                  z  |j0                        }t3        |||j4                  f||d
|\  | _        }}|g| _        t;        ||j<                  d      }|}|}d}d}g }t?        |j<                        D ]  \  }}|d	k(  r|dkD  rdnd}||k\  r|dkD  r||z  }d}||z  }|dv rdnd}g }tA        |j<                  |         D ]  }|d	k(  xr |d	k(  }t)        |j,                  |   |j.                  z  |j0                        }|tC        d0i d|d|d|jD                  dd|dz  z  d|d	k(  r|ndd|d|d|jF                  d|jH                  r|rdn|jJ                  d|j0                  d|jH                  d|jL                  d|jN                  d|d |jP                  d!|d"|d#||   |   |gz  }|d	k(  rd}||jD                  dz  z  }|}|} | xj8                  tS        ||d$| %      gz  c_        |tU        jV                  | gz  } tU        jV                  | | _,        |jZ                  rrt)        |j.                  |jZ                  z  |j0                        | _-         ||| jZ                  dfi || _.        tS        | jZ                  |d&%      | j8                  d'<   n || _-        tU        j^                         | _.         ||jZ                  d	kD  (      | _0        | jZ                  | _1        te        | jZ                  |f|| j                  d)|| _3        | ji                         D ]:  \  } }!d*| v rtk        |!tT        jl                        r|jn                  r*tT        jp                  js                  |!jt                         n+tT        jp                  jw                  |!jt                  d+d,       |!jx                  tT        jp                  js                  |!jx                         tk        |!tT        jz                        stT        jp                  j}                  |!jt                  d-d./       |!jx                  tT        jp                  js                  |!jx                         = y)1a  
        Args:
            cfg: Model architecture configuration.
            num_classes: Number of classifier classes.
            in_chans: Number of input channels.
            global_pool: Global pooling type.
            output_stride: Output stride of network, one of (8, 16, 32).
            drop_rate: Dropout rate.
            drop_path_rate: Stochastic depth drop-path rate.
            **kwargs: Extra kwargs overlayed onto cfg.
        r   Fz3Please add non-linearity constants for activation (z).)rL   )eps)rL   r   Nr   )rl   r<   T)	stagewiser   r-   rn   )r   rn   rg   rh   r$   r   r/   ri   rj   rk   r(   r0   r3   r4   r5   r9   r)   r,   r<   rl   r   zstages.r   
final_convrW   )	pool_typer   fcr~   g{Gz?fan_inlinear)modenonlinearityrF   )?rO   rP   r   r   r   grad_checkpointingr   r<   _nonlin_gammar7   r   r   r6   rd   r   r8   r   r)   r   r*   r   r'   r#   r.   r3   r   r&   r   feature_infor   r"   r   ranger|   r$   r(   r4   r0   r5   r9   r,   r   rs   r   stagesr1   r   rv   	final_acthead_hidden_sizer   headnamed_modules
isinstanceLinearr:   initzeros_weightnormal_biasConv2dkaiming_normal_)#rR   r   r   r   r   r   r   r   rq   rr   r   r   rl   r<   r)   r'   r   	stem_featdrop_path_ratesprev_chs
net_striderj   expected_varr   	stage_idxstage_depthri   rk   blocks	block_idxfirst_blockrh   nmrS   s#                                     rH   rP   zNormFreeNet.__init__  s   0 	/& ""'c$V$}}-v1dehererdssu/vv-,/,<,<(/
&s}}M#--<XYI 1A1ABJ%cmm4I =3OUXUeUefJMP^^WXcnn5IIae
!3<<#B3<<?cFVFV"VX[XbXbc,7MM-
 "-
 -
)	;	 'K3NCJJZ^_ 
&/

&; &	/"I{#q.[1_Q!F]*vzF"& J"*f"4Q!NF"3::i#89 #	'1n?a(i)@3CSCS)SUXU_U_`= #-4)) lc11 &/!^6	
 & $2  #~~ (+ww;CDTDT ::   #~~ !\\  * "mm (   *!" $39#=i#H%  ( >#%L		Q.!)"7#8 $x:X_`i_jVk"l!mmr}}f-..FM&	/N mmV, .s/?/?#BRBR/RTWT^T^ _D(43D3DaN2NDO$(1B1Bjcm$oDb! (D kkmDO"3+;+;a+?@ $ 1 1"
 "nn	

 
	 &&( 	+DAqqyZ2995##GGNN188,GGOOAHHb#666%GGNN166*Aryy)''xh'W66%GGNN166*	+rG   coarserU   c                 0    t        d|rdnddfdg      }|S )z"Group parameters for optimization.z^stemz^stages\.(\d+)z^stages\.(\d+)\.(\d+)N)z^final_conv)i )r   r   )r   )rR   r   matchers      rH   group_matcherzNormFreeNet.group_matcher
  s.     &,"2JDQ*
 rG   enablec                     || _         y)z)Enable or disable gradient checkpointing.N)r   )rR   r  s     rH   set_grad_checkpointingz"NormFreeNet.set_grad_checkpointing  s     #)rG   c                 .    | j                   j                  S )zGet the classifier head.)r   r   )rR   s    rH   get_classifierzNormFreeNet.get_classifier  s     yy||rG   c                 J    || _         | j                  j                  ||       y)zReset the classifier head.

        Args:
            num_classes: Number of classes for new classifier.
            global_pool: Global pooling type.
        N)r   r   reset)rR   r   r   s      rH   reset_classifierzNormFreeNet.reset_classifier   s     '		[1rG   rT   c                    | j                  |      }| j                  r5t        j                  j	                         st        | j                  |      }n| j                  |      }| j                  |      }| j                  |      }|S )zForward pass through feature extraction layers.

        Args:
            x: Input tensor.

        Returns:
            Feature tensor.
        )	r   r   r\   jitis_scriptingr   r   r   r   rY   s     rH   forward_featureszNormFreeNet.forward_features*  sg     IIaL""599+A+A+Ct{{A.AAAOOANN1rG   
pre_logitsc                 N    |r| j                  ||      S | j                  |      S )zForward pass through classifier head.

        Args:
            x: Input features.
            pre_logits: Return features before final linear layer.

        Returns:
            Classification logits or features.
        )r  )r   )rR   rT   r  s      rH   forward_headzNormFreeNet.forward_head<  s(     7Atyyzy2RdiiPQlRrG   c                 J    | j                  |      }| j                  |      }|S )zoForward pass.

        Args:
            x: Input tensor.

        Returns:
            Output logits.
        )r  r  rY   s     rH   rZ   zNormFreeNet.forwardH  s)     !!!$a rG   )  r   avg    r~   r~   NNrc   )T)N)r=   r>   r?   r@   r    rA   rD   rC   r   rP   r\   r
  ignorerE   r	   r   r  rs   Moduler  r
   r  r]   r  r  rZ   r^   r_   s   @rH   r   r   p  s   0  $$!#!$&B+B+ B+ 	B+
 B+ B+ B+ "B+ B+H YY	D 	T#s(^ 	 	 YY)T )T ) ) YY		  2C 2hsm 2W[ 2%,, 5<< $
Sell 
S 
S 
S %,, rG   r"   .r#   r(   r)   r*   c                 :    |xs i }t        | |ddd||||	      }|S )ar  Create NFNet ResNet configuration.

    Args:
        depths: Number of blocks in each stage.
        channels: Channel dimensions for each stage.
        group_size: Group convolution size.
        act_layer: Activation layer type.
        attn_layer: Attention layer type.
        attn_kwargs: Attention layer arguments.

    Returns:
        NFNet configuration.
    r   @   r}   )	r"   r#   r&   r'   r0   r(   r<   r)   r*   )r    )r"   r#   r(   r<   r)   r*   r   s          rH   
_nfres_cfgr  V  s:    * #K

C JrG   c                 ^    d|d   z  dz  }t        d      }t        | |dddd	|d
d|
      }|S )zCreate NFNet RegNet configuration.

    Args:
        depths: Number of blocks in each stage.
        channels: Channel dimensions for each stage.

    Returns:
        NFNet configuration.
    i   r     r/   rd_ratior%   r2   g      ?g      @Tse)
r"   r#   r&   r(   r.   r0   r1   r4   r)   r*   )r   r    )r"   r#   r1   r*   r   s        rH   
_nfreg_cfgr  z  sP     (2,&#-L$K
!C JrG   r  r0   	feat_multc                 t    t        |d   |z        }||nt        d      }t        | |dd||d||||      }	|	S )a  Create NFNet configuration.

    Args:
        depths: Number of blocks in each stage.
        channels: Channel dimensions for each stage.
        group_size: Group convolution size.
        bottle_ratio: Bottleneck ratio.
        feat_mult: Feature multiplier for final layer.
        act_layer: Activation layer type.
        attn_layer: Attention layer type.
        attn_kwargs: Attention layer arguments.

    Returns:
        NFNet configuration.
    r   r/   r  r      T)r"   r#   r&   r'   r(   r0   r5   r1   r<   r)   r*   )rA   r   r    )
r"   r#   r(   r0   r  r<   r)   r*   r1   r   s
             rH   
_nfnet_cfgr"    sZ    2 x|i/0L!,!8+dC>PK
!!C JrG   r9   c                 j    t        | |ddddddd|t        |d   dz        |dt        d      	      }|S )
a  Create DeepMind NFNet configuration.

    Args:
        depths: Number of blocks in each stage.
        channels: Channel dimensions for each stage.
        act_layer: Activation layer type.
        skipinit: Use skipinit initialization.

    Returns:
        NFNet configuration.
    r   r!  r/   Tr   r+   r  r  )r"   r#   r&   r'   r(   r0   r5   r6   r7   r9   r1   r<   r)   r*   )r    rA   r   )r"   r#   r<   r9   r   s        rH   _dm_nfnet_cfgr$    sR    " "+,#&C  JrG   dm_nfnet_f0)r   rn      r   )r"   dm_nfnet_f1)rn   r      r&  dm_nfnet_f2)r   r&     	   dm_nfnet_f3)r   r2      r(  dm_nfnet_f4)   
         dm_nfnet_f5)r&  r(  $   r*  dm_nfnet_f6)r      *      nfnet_f0nfnet_f1nfnet_f2nfnet_f3nfnet_f4nfnet_f5nfnet_f6nfnet_f7)r2      0   r-  nfnet_l0g      ?r  r}   r2   )r  
rd_divisorr;   )r"   r  r(   r0   r*   r<   eca_nfnet_l0eca)r"   r  r(   r0   r)   r*   r<   eca_nfnet_l1rn   eca_nfnet_l2eca_nfnet_l3nf_regnet_b0)r   r   r&  r&  nf_regnet_b1)rn   r   r   r   nf_regnet_b2)rn   r   r2   r2   )8   p      i  )r"   r#   nf_regnet_b3)rn   r/  r+  r+  )rM  r!     i  nf_regnet_b4)rn   r&     rS  )r        ih  nf_regnet_b5)r   r   r6  r6  )P      iP  i  nf_resnet26)rn   rn   rn   rn   nf_resnet50)r   r   r&  r   nf_resnet101)r   r      r   nf_seresnet26g      ?r  )r"   r)   r*   nf_seresnet50nf_seresnet101nf_ecaresnet26nf_ecaresnet50nf_ecaresnet101
test_nfnet)r   r   r   r   )r  r  `   r!  )r"   r#   r  r(   r0   r*   r<   variant
pretrainedr   c                 V    t         |    }t        d      }t        t        | |f||d|S )zCreate a NormFreeNet model.

    Args:
        variant: Model variant name.
        pretrained: Load pretrained weights.
        **kwargs: Additional model arguments.

    Returns:
        NormFreeNet model instance.
    T)flatten_sequential)	model_cfgfeature_cfg)
model_cfgsr   r   r   )re  rf  r   ri  rj  s        rH   _create_normfreenetrl  %  sE     7#I$/K   rG   urlc                 2    | dddddt         t        dddd	|S )
zCreate default configuration dictionary.

    Args:
        url: Model weight URL.
        **kwargs: Additional configuration options.

    Returns:
        Configuration dictionary.
    r  r      rp  r   r   ?bicubicz
stem.conv1zhead.fcz
apache-2.0)rm  r   
input_size	pool_sizecrop_pctinterpolationmeanstd
first_conv
classifierlicenser   )rm  r   s     rH   _dcfgr}  <  s4     =v)%.B")  rG   zdm_nfnet_f0.dm_in1kztimm/zmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-dnf-weights/dm_nfnet_f0-604f9c3a.pth)r&  r&  )r      r~  )r      r  rr  squash)	hf_hub_idrm  ru  rt  test_input_sizerv  	crop_modezdm_nfnet_f1.dm_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-dnf-weights/dm_nfnet_f1-fc540f82.pthrq  ro  )r   @  r  gQ?zdm_nfnet_f2.dm_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-dnf-weights/dm_nfnet_f2-89875923.pth)r2   r2   )r   `  r  gq=
ףp?zdm_nfnet_f3.dm_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-dnf-weights/dm_nfnet_f3-d74ab3aa.pth)r0  r0  )r     r  gGz?zdm_nfnet_f4.dm_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-dnf-weights/dm_nfnet_f4-0ac5b10b.pth)r(  r(  )r     r  )r      r  g;On?zdm_nfnet_f5.dm_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-dnf-weights/dm_nfnet_f5-ecb20ab1.pth)   r  )r      r  gI+?zdm_nfnet_f6.dm_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-dnf-weights/dm_nfnet_f6-e0f12116.pth)r6  r6  )r     r  )r   @  r  gd;O?)rm  ru  rt  r  )r2  r2  )r     r  )r   `  r  znfnet_l0.ra2_in1kzjhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/nfnet_l0_ra2-45c6688d.pth)r   rU  rU  )r  rm  ru  rt  r  test_crop_pctzeca_nfnet_l0.ra2_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/ecanfnet_l0_ra2-e3e9ac50.pthzeca_nfnet_l1.ra2_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/ecanfnet_l1_ra2-7dce93cd.pthzeca_nfnet_l2.ra3_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/ecanfnet_l2_ra3-da781a61.pth)rS  rS  )rm  ru  rt  r  r  r   )rm  ru  rt  r  rz  znf_regnet_b1.ra2_in1kzrhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/nf_regnet_b1_256_ra2-ad85cfef.pth)r  rm  ru  rt  r  rz  )r      r  )r     r  )r+  r+  )r     r  )rm  rz  znf_resnet50.ra2_in1kzmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/nf_resnet50_ra2-9f236009.pth)r  rm  ru  rt  r  rv  rz  )r/   r/   r/   gffffff?)r      r  )r/  r/  )r  rx  ry  rv  rt  ru  )rb  ztest_nfnet.r160_in1kc                     t        dd| i|S )z&NFNet-F0 (DeepMind weight compatible).rf  )r%  rl  rf  r   s     rH   r%  r%         NNvNNrG   c                     t        dd| i|S )z&NFNet-F1 (DeepMind weight compatible).rf  )r'  r  r  s     rH   r'  r'    r  rG   c                     t        dd| i|S )z&NFNet-F2 (DeepMind weight compatible).rf  )r)  r  r  s     rH   r)  r)    r  rG   c                     t        dd| i|S )z&NFNet-F3 (DeepMind weight compatible).rf  )r,  r  r  s     rH   r,  r,    r  rG   c                     t        dd| i|S )z&NFNet-F4 (DeepMind weight compatible).rf  )r.  r  r  s     rH   r.  r.    r  rG   c                     t        dd| i|S )z&NFNet-F5 (DeepMind weight compatible).rf  )r3  r  r  s     rH   r3  r3    r  rG   c                     t        dd| i|S )z&NFNet-F6 (DeepMind weight compatible).rf  )r5  r  r  s     rH   r5  r5    r  rG   c                     t        dd| i|S )z	NFNet-F0.rf  )r9  r  r  s     rH   r9  r9         KjKFKKrG   c                     t        dd| i|S )z	NFNet-F1.rf  )r:  r  r  s     rH   r:  r:    r  rG   c                     t        dd| i|S )z	NFNet-F2.rf  )r;  r  r  s     rH   r;  r;    r  rG   c                     t        dd| i|S )z	NFNet-F3.rf  )r<  r  r  s     rH   r<  r<    r  rG   c                     t        dd| i|S )z	NFNet-F4.rf  )r=  r  r  s     rH   r=  r=    r  rG   c                     t        dd| i|S )z	NFNet-F5.rf  )r>  r  r  s     rH   r>  r>     r  rG   c                     t        dd| i|S )z	NFNet-F6.rf  )r?  r  r  s     rH   r?  r?    r  rG   c                     t        dd| i|S )z	NFNet-F7.rf  )r@  r  r  s     rH   r@  r@    r  rG   c                     t        dd| i|S )zNFNet-L0b w/ SiLU.

    My experimental 'light' model w/ F0 repeats, 1.5x final_conv mult, 64 group_size, .25 bottleneck & SE ratio
    rf  )rC  r  r  s     rH   rC  rC    s     KjKFKKrG   c                     t        dd| i|S )zECA-NFNet-L0 w/ SiLU.

    My experimental 'light' model w/ F0 repeats, 1.5x final_conv mult, 64 group_size, .25 bottleneck & ECA attn
    rf  )rE  r  r  s     rH   rE  rE         O*OOOrG   c                     t        dd| i|S )zECA-NFNet-L1 w/ SiLU.

    My experimental 'light' model w/ F1 repeats, 2.0x final_conv mult, 64 group_size, .25 bottleneck & ECA attn
    rf  )rG  r  r  s     rH   rG  rG  $  r  rG   c                     t        dd| i|S )zECA-NFNet-L2 w/ SiLU.

    My experimental 'light' model w/ F2 repeats, 2.0x final_conv mult, 64 group_size, .25 bottleneck & ECA attn
    rf  )rH  r  r  s     rH   rH  rH  -  r  rG   c                     t        dd| i|S )zECA-NFNet-L3 w/ SiLU.

    My experimental 'light' model w/ F3 repeats, 2.0x final_conv mult, 64 group_size, .25 bottleneck & ECA attn
    rf  )rI  r  r  s     rH   rI  rI  6  r  rG   c                     t        dd| i|S )z"Normalization-Free RegNet-B0.
    rf  )rJ  r  r  s     rH   rJ  rJ  ?       O*OOOrG   c                     t        dd| i|S )z"Normalization-Free RegNet-B1.
    rf  )rK  r  r  s     rH   rK  rK  F  r  rG   c                     t        dd| i|S )z"Normalization-Free RegNet-B2.
    rf  )rL  r  r  s     rH   rL  rL  M  r  rG   c                     t        dd| i|S )z"Normalization-Free RegNet-B3.
    rf  )rP  r  r  s     rH   rP  rP  T  r  rG   c                     t        dd| i|S )z"Normalization-Free RegNet-B4.
    rf  )rR  r  r  s     rH   rR  rR  [  r  rG   c                     t        dd| i|S )z"Normalization-Free RegNet-B5.
    rf  )rV  r  r  s     rH   rV  rV  b  r  rG   c                     t        dd| i|S )z"Normalization-Free ResNet-26.
    rf  )rY  r  r  s     rH   rY  rY  i       NNvNNrG   c                     t        dd| i|S )z"Normalization-Free ResNet-50.
    rf  )rZ  r  r  s     rH   rZ  rZ  p  r  rG   c                     t        dd| i|S )z#Normalization-Free ResNet-101.
    rf  )r[  r  r  s     rH   r[  r[  w  r  rG   c                     t        dd| i|S )zNormalization-Free SE-ResNet26.rf  )r]  r  r  s     rH   r]  r]  ~       P:PPPrG   c                     t        dd| i|S )zNormalization-Free SE-ResNet50.rf  )r^  r  r  s     rH   r^  r^    r  rG   c                     t        dd| i|S )z Normalization-Free SE-ResNet101.rf  )r_  r  r  s     rH   r_  r_         QJQ&QQrG   c                     t        dd| i|S )z Normalization-Free ECA-ResNet26.rf  )r`  r  r  s     rH   r`  r`    r  rG   c                     t        dd| i|S )z Normalization-Free ECA-ResNet50.rf  )ra  r  r  s     rH   ra  ra    r  rG   c                     t        dd| i|S )z!Normalization-Free ECA-ResNet101.rf  )rb  r  r  s     rH   rb  rb    s     RZR6RRrG   c                     t        dd| i|S )z%Test NFNet model for experimentation.rf  )rc  r  r  s     rH   rc  rc    s     M
MfMMrG   )r-   )r   NNTNN))r  r  i   i   Nr[   NN))rB  h      r  )r  r     r  r!  r/   r+   r   r  N)r  r   TrF   rc   )r   )dr@   collectionsr   dataclassesr   r   	functoolsr   typingr   r   r	   r
   r   r\   torch.nnrs   	timm.datar   r   timm.layersr   r   r   r   r   r   r   r   r   r   _builderr   _features_fxr   _manipulater   	_registryr   r   __all__r    r  rJ   rD   rC   rd   rf   r|   rA   rE   r   r   r   r   r   r  r  r"  r$  rk  rl  r}  default_cfgsr%  r'  r)  r,  r.  r3  r5  r9  r:  r;  r<  r=  r>  r?  r@  rC  rE  rG  rH  rI  rJ  rK  rL  rP  rR  rV  rY  rZ  r[  r]  r^  r_  r`  ra  rb  rc  rF   rG   rH   <module>r     sT  $ $ *  7 7   A8 8 8 * 1 ' <'
"   2Eryy E8S   *'BII *'Z ABII A AN )-(,#<:<:<: <: X&	<:
 H%<: <: 2==#tCH~-.<:@ 		""	
			&c")) cP %;$($(04!c3h!S/! SM! 	!
 SM! d38n-! !HuS#X %S/ \a : %;!04(c3h(S/( ( 	(
 ( ( ( d38n-( (Z %;	!c3h!S/! ! 	!
 !H  >\2> ]3> ]3	>
 ^4> _5> _5> _5> |,> }-> }-> ~.> />  /!>" /#>$ /%>* sr$15I+>0 srdf@1>6 btdf@7>< btdf@=>B rdf@C>N <0O>P <0Q>R <:MNS>T <:MNU>V ><OPW>X ><OPY>^ ,/_>` ,/a>b =1c>f LTt]aObcg>h LTt]aObci>j ]tQU_cQdek>n \eQUQWXo>p \eQUQWXq>r mSWSYZs>v &73STcg$15Iw>
B $ # R] .s 3 4S> ( % e&5{]M\^jrte&
 5{]M\`ltve& 5{]M\`ltve& 5{}m^bnvxe&" 5{}m^cowy#e&* 5{}m^cowy+e&2 5{}m^cowy3e&< &]M[=e&@ &]M[Ae&D &]M[Ee&H (}m]Ie&L (}m]Me&P (}m]Qe&T (}m]Ue&X (}m]Ye&^ x]Madf_e&f U{]Madfge&n U{]Madfoe&v U{}mcfhwe&~ E}mcfhe&F E&]MfqsGe&J U A]M^ikKe&R E&]MfqsSe&V E&]MfqsWe&Z E(}mhsu[e&^ E(}mhsu_e&d 5RK8ee&f E{]M\`mxzge&n Eb[9oe&r Urk:se&t Urk:ue&v e{;we&z e{;{e&| e{;}e&~ <!/-6CCe& eP OD OC OK O O
 OD OC OK O O
 OD OC OK O O
 OD OC OK O O
 OD OC OK O O
 OD OC OK O O
 OD OC OK O O
 L L L L L
 L L L L L
 L L L L L
 L L L L L
 L L L L L
 L L L L L
 L L L L L
 L L L L L
 L L L L L PT PS P[ P P PT PS P[ P P PT PS P[ P P PT PS P[ P P PT PS P[ P P PT PS P[ P P PT PS P[ P P PT PS P[ P P PT PS P[ P P PT PS P[ P P OD OC OK O O OD OC OK O O PT PS P[ P P Qd Qc Qk Q Q
 Qd Qc Qk Q Q
 Rt Rs R{ R R
 Rt Rs R{ R R
 Rt Rs R{ R R
 S S S S S
 N4 N3 N; N NrG   