
    ^j                     ~   d Z ddl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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 ddlmZmZm Z m!Z! d	d
l"m#Z# d	dl$m%Z% d	dl&m'Z'm(Z( d	dl)m*Z*m+Z+m,Z, ddgZ-e G d d             Z.de/de0de0fdZ1	 dsdee0   dee/   dee0   de/deee0   ee0   f   f
dZ2	 dtde/de0de/de0de0d e0deee0   e0ee0   f   fd!Z3	 	 	 	 	 	 	 dud"e0d#e0d$e0d%e0d&e0d'eeejh                        d(e5dejh                  fd)Z6	 	 	 	 	 	 	 dud"e0d#e0d$e0d%e0d&e0d'eeejh                        d(e5dejn                  fd*Z8	 	 	 	 	 dvd+ee9   d"e0d#e0d$e0d%e0d&ee0e0f   d'eeejh                        d(e5deejh                     fd,Z: G d- d.ejh                        Z; G d/ d0ejh                        Z< G d1 d2ejh                        Z= G d3 dejh                        Z>dwd5ejh                  d6e9d7e5ddfd8Z?d9e
e9ef   de
e9ef   fd:Z@ eAdxi d; e.d<d=d>dd?@      dA e.d<dBdCdDdE@      dF e.d<dBdCdDdEdGH      dI e.dJdKdLd<dD@      dM e.dNdOdPdDdD@      dQ e.dRdSdTd<dU@      dV e.dWdXdTdJdY@      dZ e.d[d\d]d^d_@      d` e.dadbdcdNdd@      de e.dRdfdgdhd_@      di e.djdkdldmdn@      do e.dpdqdrdsdE@      dt e.dudvdwdjd_@      dx e.d<d=d>dd?dyz      d{ e.dJd|d}ddDdyz      d~ e.dJdddDddyz      d e.dNdddDddyz      d e.dNdddDddydG      d e.dJddd<ddyz      d e.dRddd<ddyz      d e.d[ddLddEdyz      d e.dmddddYdyz      d e.ddddNdddyz      d e.ddddNdddydG      d e.djdkdldmdndyz      d e.ddddmdUdyz      d e.ddddddyz      d e.ddddddyz      d e.ddddddyz      d e.ddddddyz      d e.d[ddLddEdyd eedD            d e.dEd[ddLddydd      d e.dYdmddddyddd	      d e.ddDdddddyddddë      d e.ddJdddddyddddë      d e.ddJdddddyddddë      ZBde9de5de>fd̄ZCdyde9de
e9ef   fd΄ZDdyde9de
e9ef   fdτZEdyde9de
e9ef   fdЄZF e*i d eDddӬԫ      d eDdd֬ԫ      d eDddجԫ      d eDddڬԫ      d eDdҬܫ      d eDdҬܫ      d eDdҬܫ      d eDdd      d eDdd      d eDd4      d eDddd      d eDddd      d eDd4      d eDdddddd      d eDdddddd      d eDddԫ      d eFddԫ      i d eFddԫ      d eFddԫ      d eFddԫ      d eFdҐd ԫ      d eFdҐdԫ      d eFdҐdԫ      d eFdҐdԫ      d eFdҐdԫ      d	 eFdҐd
ԫ      d eFdҐdԫ      d eFdҐdԫ      d eFdҐdԫ      d eFdҐdԫ      d eFdҐddddd      d eFdҐddddd      d eFdҐddddd      d eFdҐdd      i d  eFdҐd!d      d" eFdҐd#d      d$ eFdҐd%d&d'ddd(      d) eFdҐd%d&d*ddd(      d+ eFdҐd%d&d,ddd(      d- eFdҐd%d&d.ddd(      d/ eFdҐd0dd%d&1      d2 eFdҐd3dd%d&1      d4 eFdҐd5dd%d&1      d6 eEdҬܫ      d7 eEdҬܫ      d8 eEdҬܫ      d9 eEdҬܫ      d: eEdҬܫ      d; eEdҬܫ      d< eEdҬܫ      d= eEdҬܫ      i d> eEdҬܫ      d? eEdҬܫ      d@ eEdҬܫ      dA eEdҬܫ      dB eEdҬܫ      dC eEdҬܫ      dD eEdҬܫ      dE eEdҬܫ      dF eEdҬܫ      dG eEdҬܫ      dH eEdҬܫ      dI eEdҬܫ      dJ eEdҬܫ      dK eEdҬܫ      dL eEdҬܫ      dM eEdҬܫ            ZGe+dzde5de>fdN       ZHe+dzde5de>fdO       ZIe+dzde5de>fdP       ZJe+dzde5de>fdQ       ZKe+dzde5de>fdR       ZLe+dzde5de>fdS       ZMe+dzde5de>fdT       ZNe+dzde5de>fdU       ZOe+dzde5de>fdV       ZPe+dzde5de>fdW       ZQe+dzde5de>fdX       ZRe+dzde5de>fdY       ZSe+dzde5de>fdZ       ZTe+dzde5de>fd[       ZUe+dzde5de>fd\       ZVe+dzde5de>fd]       ZWe+dzde5de>fd^       ZXe+dzde5de>fd_       ZYe+dzde5de>fd`       ZZe+dzde5de>fda       Z[e+dzde5de>fdb       Z\e+dzde5de>fdc       Z]e+dzde5de>fdd       Z^e+dzde5de>fde       Z_e+dzde5de>fdf       Z`e+dzde5de>fdg       Zae+dzde5de>fdh       Zbe+dzde5de>fdi       Zce+dzde5de>fdj       Zde+dzde5de>fdk       Zee+dzde5de>fdl       Zfe+dzde5de>fdm       Zge+dzde5de>fdn       Zhe+dzde5de>fdo       Zie+dzde5de>fdp       Zje+dzde5de>fdq       Zk e,eldrdi       y({  a  RegNet X, Y, Z, and more

Paper: `Designing Network Design Spaces` - https://arxiv.org/abs/2003.13678
Original Impl: https://github.com/facebookresearch/pycls/blob/master/pycls/models/regnet.py

Paper: `Fast and Accurate Model Scaling` - https://arxiv.org/abs/2103.06877
Original Impl: None

Based on original PyTorch impl linked above, but re-wrote to use my own blocks (adapted from ResNet here)
and cleaned up with more descriptive variable names.

Weights from original pycls impl have been modified:
* first layer from BGR -> RGB as most PyTorch models are
* removed training specific dict entries from checkpoints and keep model state_dict only
* remap names to match the ones here

Supports weight loading from torchvision and classy-vision (incl VISSL SEER)

A number of custom timm model definitions additions including:
* stochastic depth, gradient checkpointing, layer-decay, configurable dilation
* a pre-activation 'V' variant
* only known RegNet-Z model definitions with pretrained weights

Hacked together by / Copyright 2020 Ross Wightman
    N)	dataclassreplace)partial)AnyCallableDictListOptionalUnionTupleTypeIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)ClassifierHeadAvgPool2dSameConvNormActSEModuleDropPathGroupNormActcalculate_drop_path_rates)get_act_layerget_norm_act_layercreate_conv2dmake_divisible   )build_model_with_cfg)feature_take_indices)checkpoint_seqnamed_apply)generate_default_cfgsregister_modelregister_model_deprecationsRegNet	RegNetCfgc                       e Zd ZU dZdZeed<   dZeed<   dZe	ed<   dZ
e	ed	<   d
Zeed<   dZe	ed<   dZe	ed<   dZe	ed<   dZeed<   dZee   ed<   dZeed<   dZeed<   dZeed<   dZeeef   ed<   dZeeef   ed<   y)r%   z"RegNet architecture configuration.   depthP   w0q=
ףPE@waHzG@wm   
group_size      ?bottle_ratio        se_ratiogroup_min_ratio    
stem_widthconv1x1
downsampleF
linear_outpreactr   num_featuresrelu	act_layer	batchnorm
norm_layerN)__name__
__module____qualname____doc__r(   int__annotations__r*   r,   floatr.   r0   r2   r4   r5   r7   r9   r
   strr:   boolr;   r<   r>   r   r   r@        ]/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/regnet.pyr%   r%   -   s    ,E3OBLBBJL%HeOUJ )J)JFDL#&,IuS(]#,'2Jc8m$2rK   fqreturnc                 6    t        t        | |z        |z        S )zConverts a float to the closest non-zero int divisible by q.

    Args:
        f: Input float value.
        q: Quantization divisor.

    Returns:
        Quantized integer value.
    )rE   round)rM   rN   s     rL   quantize_floatrR   A   s     uQU|a  rK   widthsbottle_ratiosgroups	min_ratioc           	         t        | |      D cg c]  \  }}t        ||z         }}}t        ||      D cg c]  \  }}t        ||       }}}|r*t        ||      D cg c]  \  }}t        |||       }}}n(t        ||      D cg c]  \  }}t	        ||       }}}t        ||      D cg c]  \  }}t        ||z         } }}| |fS c c}}w c c}}w c c}}w c c}}w c c}}w )a,  Adjusts the compatibility of widths and groups.

    Args:
        widths: List of channel widths.
        bottle_ratios: List of bottleneck ratios.
        groups: List of group sizes.
        min_ratio: Minimum ratio for divisibility.

    Returns:
        Tuple of adjusted widths and groups.
    )ziprE   minr   rR   )	rS   rT   rU   rV   wbbottleneck_widthsgw_bots	            rL   adjust_widths_groups_compr_   N   s    " 14FM0JK1QUKK,/8I,JK5c!UmKFKQTUfhnQopXUA^E1i@ppFIJ[]cFde(%^E15ee-01BM-RSc%!)nSFS6> LK qeSs   CC$C!C'6C-   width_slopewidth_initial
width_multr(   r0   quantc                    | dk\  r|dkD  r|dkD  r||z  dk(  sJ t        j                  |t         j                        | z  |z   }t        j                  t        j                  ||z        t        j                  |      z        }t        j                  |t        j                  ||      z  |z        |z  }t        t        j                  |            t        |j                         j                               dz   }
}	t        j                  t        |	      D cg c]  }| c}t         j                        }|j                         j                         |	|j                         fS c c}w )au  Generates per block widths from RegNet parameters.

    Args:
        width_slope: Slope parameter for width progression.
        width_initial: Initial width.
        width_mult: Width multiplier.
        depth: Network depth.
        group_size: Group convolution size.
        quant: Quantization factor.

    Returns:
        Tuple of (widths, num_stages, groups).
    r   r   )dtype)torcharangefloat32rQ   logmathpowlenuniquerE   maxitemtensorrangeint32tolist)ra   rb   rc   r(   r0   rd   widths_cont
width_expsrS   
num_stages	max_stage_rU   s                rL   generate_regnetrz   j   s   * ! 1j1nY^I^bcIccc ,,uEMM:[H=XKUYY{]'BCdhhzFZZ[J[[-%))J
*KKuTUX]]FV 45s:>>;K;P;P;R7SVW7W	J\\uZ/@A!:AUF::< *fmmo== Bs   	E+in_chsout_chskernel_sizestridedilationr@   r;   c	                     ||d}	|xs t         j                  }|dk(  r|dk(  rdn|}|dkD  r|nd}|rt        | ||f||d|	S t        | ||f|||dd|	S )am  Create convolutional downsampling module.

    Args:
        in_chs: Input channels.
        out_chs: Output channels.
        kernel_size: Convolution kernel size.
        stride: Convolution stride.
        dilation: Convolution dilation.
        norm_layer: Normalization layer.
        preact: Use pre-activation.

    Returns:
        Downsampling module.
    devicerf   r   )r~   r   F)r~   r   r@   	apply_act)nnBatchNorm2dr   r   )
r{   r|   r}   r~   r   r@   r;   r   rf   dds
             rL   downsample_convr      s    2 U	+B-r~~J{x1}!+K&?xH
 
 
 	
 	
 !	
 	
 		
rK   c	                 L   ||d}	|xs t         j                  }|dk(  r|nd}
t        j                         }|dkD  s|dkD  r,|
dk(  r|dkD  rt        nt         j                  } |d|
dd      }|rt        | |dfddi|	}nt        | |dfd|dd|	}t        j                  ||g S )	a  Create average pool downsampling module.

    AvgPool Downsampling as in 'D' ResNet variants. This is not in RegNet space but I might experiment.

    Args:
        in_chs: Input channels.
        out_chs: Output channels.
        kernel_size: Convolution kernel size.
        stride: Convolution stride.
        dilation: Convolution dilation.
        norm_layer: Normalization layer.
        preact: Use pre-activation.

    Returns:
        Sequential downsampling module.
    r   r      TF)	ceil_modecount_include_padr~   )r~   r@   r   )r   r   Identityr   	AvgPool2dr   r   
Sequential)r{   r|   r}   r~   r   r@   r;   r   rf   r   
avg_stridepoolavg_pool_fnconvs                 rL   downsample_avgr      s    6 U	+B-r~~J#q=aJ;;=DzX\'1Q8a<mR\\1jDERVWa@@R@67AfaJZ_fcef==4,''rK   downsample_typec
                     ||	d}
| dv sJ ||k7  s|dk7  s|d   |d   k7  r7t        d	||d   ||d|
}| sy| dk(  rt        ||fi |S t        ||fd|i|S t        j                         S )
a  Create shortcut connection for residual blocks.

    Args:
        downsample_type: Type of downsampling ('avg', 'conv1x1', or None).
        in_chs: Input channels.
        out_chs: Output channels.
        kernel_size: Kernel size for conv downsampling.
        stride: Stride for downsampling.
        dilation: Dilation rates.
        norm_layer: Normalization layer.
        preact: Use pre-activation.

    Returns:
        Shortcut module or None.
    r   )avgr8    Nr   r   )r~   r   r@   r;   Nr   r}   rJ   )dictr   r   r   r   )r   r{   r|   r}   r~   r   r@   r;   r   rf   r   dargss               rL   create_shortcutr      s    6 U	+B::::FaK8A;(1++EeFXa[ZX^ebde%!&';U;;"67UUuUU{{}rK   c                   :    e Zd ZdZdddddddej
                  ej                  ddddfd	ed
ededeeef   de	dede	de
dedeej                     deej                     deeej                        de	f fdZddZdej$                  dej$                  fdZ xZS )
BottleneckzRegNet Bottleneck block.

    This is almost exactly the same as a ResNet Bottleneck. The main difference is the SE block is moved from
    after conv3 to after conv2. Otherwise, it's just redefining the arguments for groups/bottleneck channels.
    r   r   r         ?r8   FNr3   r{   r|   r~   r   r2   r0   r4   r9   r:   r>   r@   
drop_blockdrop_path_ratec           	         ||d}t         |           t        |
      }
t        t	        ||z              }||z  }t        |
|      }t        ||fddi||| _        t        ||fd||d   ||d||| _        |r,t        t	        ||z              }t        |f||
d|| _
        nt        j                         | _
        t        ||fdd	d
||| _        |	rt        j                         n |
       | _        t        |||fd|||d|| _        |dkD  rt#        |      | _        yt        j                         | _        y)a  Initialize RegNet Bottleneck block.

        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            stride: Convolution stride.
            dilation: Dilation rates for conv2 and shortcut.
            bottle_ratio: Bottleneck ratio (reduction factor).
            group_size: Group convolution size.
            se_ratio: Squeeze-and-excitation ratio.
            downsample: Shortcut downsampling type.
            linear_out: Use linear activation for output.
            act_layer: Activation layer.
            norm_layer: Normalization layer.
            drop_block: Drop block layer.
            drop_path_rate: Stochastic depth drop rate.
        r   r>   r@   r}   r      r   )r}   r~   r   rU   
drop_layerrd_channelsr>   F)r}   r   )r}   r~   r   r@   N)super__init__r   rE   rQ   r   r   conv1conv2r   ser   r   conv3act3r   r9   r   	drop_path)selfr{   r|   r~   r   r2   r0   r4   r9   r:   r>   r@   r   r   r   rf   r   bottleneck_chsrU   cargsse_channels	__class__s                        rL   r   zBottleneck.__init__  sv   F /!),	U7\#9:;:-yZ@ VQV%VSUV
 

 a[!

 

 


 eFX$567K~b;R[b_abDGkkmDG haSXh\ahegh
%/BKKMY[	)	
 !	
 	
 6Da5G.1R[[]rK   rO   c                 ~    t         j                  j                  | j                  j                  j
                         y)z1Zero-initialize the last batch norm in the block.N)r   initzeros_r   bnweightr   s    rL   zero_init_lastzBottleneck.zero_init_last`  s     
tzz}}++,rK   xc                    |}| j                  |      }| j                  |      }| j                  |      }| j                  |      }| j                  #| j                  |      | j	                  |      z   }| j                  |      }|S zoForward pass.

        Args:
            x: Input tensor.

        Returns:
            Output tensor.
        )r   r   r   r   r9   r   r   r   r   shortcuts      rL   forwardzBottleneck.forwardd  sw     JJqMJJqMGGAJJJqM??& q!DOOH$==AIIaLrK   rO   NrA   rB   rC   rD   r   ReLUr   rE   r   rG   rH   rI   r   Moduler
   r   r   rg   Tensorr   __classcell__r   s   @rL   r   r     s     (."#"'$)+*,..48$&!G[G[ G[ 	G[
 CHoG[  G[ G[ G[ G[ G[ BIIG[ RYYG[ !bii1G[ "G[R- %,, rK   r   c                   :    e Zd ZdZdddddddej
                  ej                  ddddfd	ed
ededeeef   de	dede	de
dedeej                     deej                     deeej                        de	f fdZddZdej$                  dej$                  fdZ xZS )PreBottleneckznPre-activation RegNet Bottleneck block.

    Similar to Bottleneck but with pre-activation normalization.
    r   r   r   r8   FNr3   r{   r|   r~   r   r2   r0   r4   r9   r:   r>   r@   r   r   c                 p   ||d}t         |           t        ||
      }t        t	        ||z              }||z  } ||fi || _        t        ||fddi|| _         ||fi || _        t        ||fd||d   |d|| _	        |r,t        t	        ||z              }t        |f||
d|| _        nt        j                         | _         ||fi || _        t        ||fddi|| _        t!        |||fd||dd	|| _        |dkD  rt%        |      | _        y
t        j                         | _        y
)a  Initialize pre-activation RegNet Bottleneck block.

        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            stride: Convolution stride.
            dilation: Dilation rates for conv2 and shortcut.
            bottle_ratio: Bottleneck ratio (reduction factor).
            group_size: Group convolution size.
            se_ratio: Squeeze-and-excitation ratio.
            downsample: Shortcut downsampling type.
            linear_out: Use linear activation for output.
            act_layer: Activation layer.
            norm_layer: Normalization layer.
            drop_block: Drop block layer.
            drop_path_rate: Stochastic depth drop rate.
        r   r}   r   r   r   )r}   r~   r   rU   r   T)r}   r~   r   r;   N)r   r   r   rE   rQ   norm1r   r   norm2r   r   r   r   r   norm3r   r   r9   r   r   )r   r{   r|   r~   r   r2   r0   r4   r9   r:   r>   r@   r   r   r   rf   r   norm_act_layerr   rU   r   r   s                        rL   r   zPreBottleneck.__init__  so   F /+J	BU7\#9:;:-#F1b1
"6>OqOBO
#N9b9
"
 a[
 

 eFX$567K~b;R[b_abDGkkmDG#N9b9
">7PPRP
)	
 	
 	
 6Da5G.1R[[]rK   rO   c                      y)z?Zero-initialize the last batch norm (no-op for pre-activation).NrJ   r   s    rL   r   zPreBottleneck.zero_init_last  s    rK   r   c                 V   | j                  |      }|}| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }| j                  |      }| j                  #| j                  |      | j                  |      z   }|S r   )	r   r   r   r   r   r   r   r9   r   r   s      rL   r   zPreBottleneck.forward  s     JJqMJJqMJJqMJJqMGGAJJJqMJJqM??& q!DOOH$==ArK   r   r   r   s   @rL   r   r   z  s     (."#"'$)+*,..48$&!F[F[ F[ 	F[
 CHoF[  F[ F[ F[ F[ F[ BIIF[ RYYF[ !bii1F[ "F[P %,, rK   r   c                        e Zd ZdZdefdedededededeee      d	e	e
j                     f fd
Zdej                  dej                  fdZ xZS )RegStagezRegNet stage (sequence of blocks with the same output shape).

    A stage consists of multiple bottleneck blocks with the same output dimensions.
    Nr(   r{   r|   r~   r   drop_path_ratesblock_fnc                    t         |           d| _        |dv rdnd}	t        |      D ]U  }
|
dk(  r|nd}|
dk(  r|n|}|	|f}|||
   nd}dj	                  |
dz         }| j                  | |||f|||d	|       |}	W y)
a  Initialize RegNet stage.

        Args:
            depth: Number of blocks in stage.
            in_chs: Input channels.
            out_chs: Output channels.
            stride: Stride for first block.
            dilation: Dilation rate.
            drop_path_rates: Drop path rates for each block.
            block_fn: Block class to use.
            **block_kwargs: Additional block arguments.
        F)r   r   r   r   r   Nr3   zb{})r~   r   r   )r   r   grad_checkpointingrr   format
add_module)r   r(   r{   r|   r~   r   r   r   block_kwargsfirst_dilationiblock_strideblock_in_chsblock_dilationdprnamer   s                   rL   r   zRegStage.__init__  s    . 	"'&&0au 	&A%&!V6L%&!V6L,h7N(7(C/!$C<<A&DOO  (+#& #
 &N#	&rK   r   rO   c                     | j                   r:t        j                  j                         st	        | j                         |      }|S | j                         D ]
  } ||      } |S )zForward pass through all blocks in the stage.

        Args:
            x: Input tensor.

        Returns:
            Output tensor.
        )r   rg   jitis_scriptingr   children)r   r   blocks      rL   r   zRegStage.forward  sZ     ""599+A+A+Ct}}2A   !HrK   )rA   rB   rC   rD   r   rE   r
   r	   rG   r   r   r   r   rg   r   r   r   r   s   @rL   r   r     s     6:(2,&,& ,& 	,&
 ,& ,& &d5k2,& 299o,&\ %,, rK   r   c                   :    e Zd ZdZ	 	 	 	 	 	 	 	 	 d#dedededededed	ed
ef fdZ		 	 	 d$dededed	ede
eeeef      eeef   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(dej2                  deeeee   f      dededededeeej2                     e
ej2                  eej2                     f   f   fdZ	 	 	 d)deeee   f   dededee   fdZdej2                  dej2                  fdZd%dej2                  d edej2                  fd!Zdej2                  dej2                  fd"Z xZ S )*r$   zRegNet-X, Y, and Z Models.

    Paper: https://arxiv.org/abs/2003.13678
    Original Impl: https://github.com/facebookresearch/pycls/blob/master/pycls/models/regnet.py
    Ncfgin_chansnum_classesoutput_strideglobal_pool	drop_rater   r   c           
         t         |           |	|
d}|| _        || _        || _        |dv sJ t        |fi |}|j                  }t        |j                  |j                        }|j                  rt        ||dfddi|| _        nt        ||dfddi||| _        t        |dd      g| _        |}d}| j                  |||	      \  }}t!        |      d
k(  sJ |j                  rt"        nt$        }t'        |      D ]j  \  }}dj)                  |dz         }| j+                  |t-        d||d|||       |d   }||d   z  }| xj                  t        |||      gz  c_        l |j.                  r2t        ||j.                  fddi||| _        |j.                  | _        nV|j2                  xs |j                  }|r t5        |j                               nt7        j8                         | _        || _        | j.                  | _        t=        d| j.                  |||d|| _        tA        tC        tD        |      |        y)a  Initialize RegNet model.

        Args:
            cfg: Model architecture configuration.
            in_chans: Number of input channels.
            num_classes: Number of classifier classes.
            output_stride: Output stride of network, one of (8, 16, 32).
            global_pool: Global pooling type.
            drop_rate: Dropout rate.
            drop_path_rate: Stochastic depth drop-path rate.
            zero_init_last: Zero-init last weight of residual path.
            kwargs: Extra kwargs overlayed onto cfg.
        r   )r`      r6   r   r   r~   r   stem)num_chs	reductionmodule)r   r      zs{}r   )r{   r   r|   r}   )in_featuresr   	pool_typer   )r   NrJ   )#r   r   r   r   r   r   r7   r   r>   r@   r;   r   r   r   feature_info_get_stage_argsrm   r   r   	enumerater   r   r   r<   
final_convr:   r   r   r   head_hidden_sizer   headr    r   _init_weights)r   r   r   r   r   r   r   r   r   r   rf   kwargsr   r7   na_args
prev_widthcurr_strideper_stage_argscommon_argsr   r   
stage_args
stage_name	final_actr   s                           rL   r   zRegNet.__init__0  s~   6 	/& "+++c$V$ ^^
3>>J::%h
ANaN2NDI#Hj!WAWWTVWDI!*&QR  
&*&:&:') '; '
#
 >"a'''$'JJ=J&~6 	fMAza!e,JOO %% ! "	
 	 $I.J:h//K$z[Yc"d!ee	f" )*c6F6FgTUgY`gdfgDO # 0 0D4#**I@I:mCMM:<r{{}DO *D $ 1 1" 
))#!	

 
	 	GM.I4PrK   default_striderO   c           	      d   t        |j                  |j                  |j                  |j                  |j
                        \  }}}t        j                  t        j                  |      d      \  }}	|j                         |	j                         }	}t        |      D 
cg c]  }
|j                   }}
g }g }d}d}t        |      D ]8  }
||k\  r||z  }d}n|}||z  }|j                  |       |j                  |       : t        ||	d      }t        ||||j                        \  }}g d}t!        ||||	|||      D cg c]  }t#        t!        ||             }}t#        |j$                  |j&                  |j(                  |j*                  |j,                        }||fS c c}
w c c}w )	aM  Generate stage arguments from configuration.

        Args:`
            cfg: RegNet configuration.
            default_stride: Default stride for stages.
            output_stride: Target output stride.
            drop_path_rate: Stochastic depth rate.

        Returns:
            Tuple of (per_stage_args, common_args).
        T)return_countsr   r   )	stagewise)rV   )r|   r~   r   r(   r2   r0   r   )r9   r4   r:   r>   r@   )rz   r,   r*   r.   r(   r0   rg   rn   rq   rt   rr   r2   appendr   r_   r5   rX   r   r9   r4   r:   r>   r@   )r   r   r  r   r   rS   rw   stage_gsstage_widthsstage_depthsry   stage_brstage_stridesstage_dilations
net_strider   r~   	stage_dpr	arg_namesparamsr  r  s                         rL   r   zRegNet._get_stage_args  s   & (7svvsvvsvvsyyZ]ZhZh'i$
H &+\\%,,v2FVZ%["l%1%8%8%:L<O<O<Ql.3J.?@C$$@@
z" 	-A]*N*'f$
  (""8,	- .nlVZ[	!:(H8K8K"Mho	 m_lHV^`ij
-3DY'(
 
 ~~\\~~mm~~
 {**= A&
s    F(F-coarsec                 .    t        d|rd      S d      S )z"Group parameters for optimization.z^stemz^s(\d+)z^s(\d+)\.b(\d+))r   blocks)r   )r   r  s     rL   group_matcherzRegNet.group_matcher  s&     !':
 	
-?
 	
rK   enablec                 T    t        | j                               dd D ]	  }||_         y)z)Enable or disable gradient checkpointing.r   N)listr   r   )r   r  ss      rL   set_grad_checkpointingzRegNet.set_grad_checkpointing  s-     dmmo&q, 	*A#)A 	*rK   c                 .    | j                   j                  S )zGet the classifier head.)r   fcr   s    rL   get_classifierzRegNet.get_classifier  s     yy||rK   c                 L    || _         | j                  j                  ||       y)zReset the classifier head.

        Args:
            num_classes: Number of classes for new classifier.
            global_pool: Global pooling type.
        )r   N)r   r   reset)r   r   r   s      rL   reset_classifierzRegNet.reset_classifier  s      '		{;rK   r   indicesnorm
stop_early
output_fmtintermediates_onlyc                 @   |dv sJ d       g }t        d|      \  }}	d}
| j                  |      }|
|v r|j                  |       d}|r|d|	 }|D ]/  }|
dz  }
 t        | |      |      }|
|v s|j                  |       1 |r|S |
d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   s1s2s3s4Nr   r   )r   r   r
  getattrr   )r   r   r%  r&  r'  r(  r)  intermediatestake_indices	max_indexfeat_idxlayer_namesns                rL   forward_intermediateszRegNet.forward_intermediates  s    * Y&D(DD&"6q'"Bi IIaL|#  #.%jy1K 	(AMH a #A<'$$Q'		(   q="A-rK   
prune_norm
prune_headc                     t        d|      \  }}d}||d }|D ]!  }t        | |t        j                                # |dk  rt        j                         | _        |r| j                  dd       |S )aE  Prune layers not required for specified intermediates.

        Args:
            indices: Indices of intermediate layers to keep.
            prune_norm: Whether to prune normalization layer.
            prune_head: Whether to prune the classifier head.

        Returns:
            List of indices that were kept.
        r,  r-  Nr   r   r   )r   setattrr   r   r   r$  )r   r%  r:  r;  r4  r5  r7  r8  s           rL   prune_intermediate_layersz RegNet.prune_intermediate_layers  st      #7q'"Bi.!)*- 	,AD!R[[]+	,q= kkmDO!!!R(rK   c                     | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }|S )zForward pass through feature extraction layers.

        Args:
            x: Input tensor.

        Returns:
            Feature tensor.
        )r   r.  r/  r0  r1  r   r   r   s     rL   forward_featureszRegNet.forward_features)  sX     IIaLGGAJGGAJGGAJGGAJOOArK   
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.
        )rB  )r   )r   r   rB  s      rL   forward_headzRegNet.forward_head:  s(     7Atyyzy2RdiiPQlRrK   c                 J    | j                  |      }| j                  |      }|S )zoForward pass.

        Args:
            x: Input tensor.

        Returns:
            Output logits.
        )rA  rD  r@  s     rL   r   zRegNet.forwardF  s)     !!!$a rK   )	r     r6   r   r3   r3   TNN)r   r6   r3   F)T)N)NFFr+  F)r   FT)!rA   rB   rC   rD   r%   rE   rH   rG   rI   r   r   r	   r   r   r   rg   r   ignorer  r  r   r   r!  r
   r$  r   r   r9  r>  rA  rD  r   r   r   s   @rL   r$   r$   )  s    #!#$!$&#'WQWQ WQ 	WQ
 WQ WQ WQ "WQ !WQx #$!#$&6+6+  6+ 	6+
 "6+ 
tDcN#T#s(^3	46+p YY
D 
T#s(^ 
 
 YY*T *T * *
 YY		  <C <hsm <W[ < 8<$$',. ||.  eCcN34.  	. 
 .  .  !%.  
tELL!5tELL7I)I#JJ	K. d ./$#	3S	>*  	
 
c6%,, 5<< "
Sell 
S 
S 
S %,, rK   r   r   r   r   c                    t        | t        j                        r| j                  d   | j                  d   z  | j                  z  }|| j
                  z  }| j                  j                  j                  dt        j                  d|z               | j                  %| j                  j                  j                          yyt        | t        j                        rct        j                  j                  | j                  dd       | j                  *t        j                  j                  | j                         yy|rt!        | d      r| j#                          yyy)	zInitialize module weights.

    Args:
        module: PyTorch module to initialize.
        name: Module name.
        zero_init_last: Zero-initialize last layer weights.
    r   r          @Nr3   g{Gz?)meanstdr   )
isinstancer   Conv2dr}   out_channelsrU   r   datanormal_rk   sqrtbiaszero_Linearr   r   hasattrr   )r   r   r   fan_outs       rL   r   r   T  s    &"))$$$Q'&*<*<Q*??&BUBUUFMM!""1diig&>?;;"KK""$ #	FBII	&
CT:;;"GGNN6;;' #	GF,<= >rK   
state_dictc                    | j                  d|       } g d}d| v rddl}| d   d   d   } i }| d   j                         D ]q  \  }}|j                  dd	      }|j                  d
d      }|j	                  dd |      }|j	                  dd|      }|D ]  \  }}|j                  ||      } |||<   s | d   j                         D ]%  \  }}d|v sd|v r|j                  dd      }|||<   ' |S d| v rddl}i }| j                         D ]p  \  }}|j                  dd	      }|j                  dd      }|j	                  dd |      }|D ]  \  }}|j                  ||      } |j                  dd      }|||<   r |S | S )zFilter and remap state dict keys for compatibility.

    Args:
        state_dict: Raw state dictionary.

    Returns:
        Filtered state dictionary.
    model))zf.a.0z
conv1.conv)zf.a.1zconv1.bn)zf.b.0z
conv2.conv)zf.b.1zconv2.bn)z
f.final_bnconv3.bn)zf.se.excitation.0zse.fc1)zf.se.excitation.2zse.fc2)zf.ser   )zf.c.0
conv3.conv)zf.c.1r[  )zf.cr\  )zproj.0downsample.conv)zproj.1zdownsample.bn)projr]  classy_state_dictr   N
base_modeltrunkz_feature_blocks.conv1.stem.0	stem.convz_feature_blocks.conv1.stem.1zstem.bnz&^_feature_blocks.res\d.block(\d)-(\d+)c                 x    dt        | j                  d             dt        | j                  d            dz    S )Nr  r   .br   rE   groupr   s    rL   <lambda>z_filter_fn.<locals>.<lambda>  2    Ac!''!*o.bQWWQZ11D0EF rK   zs(\d)\.b(\d+)\.bnzs\1.b\2.downsample.bnheadsprojection_head
prototypesz0.clf.0head.fczstem.0.weightzstem.0zstem.1z)trunk_output.block(\d)\.block(\d+)\-(\d+)c                 x    dt        | j                  d             dt        | j                  d            dz    S )Nr  r   rd  r   re  rg  s    rL   rh  z_filter_fn.<locals>.<lambda>  ri  rK   zfc.zhead.fc.)getreitemsr   sub)rX  replacesrp  outkvr  rs           rL   
_filter_fnrx  j  s    4JH  j( 34\B7K
w'--/ 		DAq		8+FA		8)DA9FKA +-EqIA  $1IIaO$CF		 w'--/ 	DAq A%):		)Y/ACF		
 
*$$$& 		DAq		(K0A		(I.A<FKA ! $1IIaO$		%,ACF		 
rK   regnetx_002r/   gQ8B@gQ@   )r*   r,   r.   r0   r(   regnetx_004g{Gz8@gRQ@r      regnetx_004_tvg?)r*   r,   r.   r0   r(   r5   regnetx_0060   g\(|B@gQ@regnetx_0088   g=
ףpA@g=
ףp=@regnetx_016r)   gzGA@g      @   regnetx_032X   g(\O:@   regnetx_040`   g33333SC@gq=
ףp@(      regnetx_064   g
ףp=jN@g(\ @   regnetx_080gHzH@g
ףp=
@x   regnetx_120   gףp=
WR@g(\@p      regnetx_160   gQK@g @   regnetx_320@  gףp=
wQ@rJ  regnety_002r   )r*   r,   r.   r0   r(   r4   regnety_004gp=
;@gQ @regnety_006gQE@@g(\@   regnety_008gQkC@g333333@   regnety_008_tv)r*   r,   r.   r0   r(   r4   r5   regnety_016g(\µ4@g333333@   regnety_032r+   r-   r'   regnety_040g)\h?@@   regnety_064g\(@@g)\(@H   regnety_080   gGz4S@gQ@regnety_080_tvregnety_120regnety_160   gQZ@gףp=
@regnety_320   g)\\@g=
ףp=@   regnety_640i`  g(\ob@iH  regnety_1280i  g(\d@g)\(@i  regnety_2560i  g(\l@iu  regnety_040_sgnsilu)r0   )r*   r,   r.   r0   r(   r4   r>   r@   regnetv_040T)r(   r*   r,   r.   r0   r4   r;   r>   regnetv_064r   )	r(   r*   r,   r.   r0   r4   r;   r>   r9   regnetz_005gffffff%@gGz@r   g      @i   )r(   r*   r,   r.   r0   r2   r4   r9   r:   r<   r>   regnetz_040   g      -@g+@regnetz_040_hi   variant
pretrainedc                 B    t        t        | |ft        |    t        d|S )zCreate a RegNet model.

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

    Returns:
        RegNet model instance.
    )	model_cfgpretrained_filter_fn)r   r$   
model_cfgsrx  )r  r  r   s      rL   _create_regnetr    s2      W%' 	 rK   urlc                 6    | dd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.
    rF  r      r     r  )r      r  gffffff?r1   bicubicrb  rm  z
apache-2.0)r  r   
input_size	pool_sizetest_input_sizecrop_pcttest_crop_pctinterpolationrK  rL  
first_conv
classifierlicenser   r  r   s     rL   _cfgr    s:     4}SY(dS",AJ^!
 $* rK   c                 4    | dddddt         t        dddd	d
|S )zCreate pycls configuration dictionary.

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

    Returns:
        Configuration dictionary.
    rF  r  r  g      ?r  rb  rm  mitz)https://github.com/facebookresearch/pyclsr  r   r  r  r  r  rK  rL  r  r  r  
origin_urlr   r  s     rL   _cfgpycr    s:     4}SYI%.B!(S
 X^ rK   c                 4    | dddddt         t        dddd	d
|S )zCreate torchvision v2 configuration dictionary.

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

    Returns:
        Configuration dictionary.
    rF  r  r  gzG?r  rb  rm  zbsd-3-clausez!https://github.com/pytorch/visionr  r   r  s     rL   _cfgtv2r  $  s:     4}SYI%.B!!1T
 Y_ rK   zregnety_032.ra_in1kztimm/znhttps://github.com/huggingface/pytorch-image-models/releases/download/v0.1-weights/regnety_032_ra-7f2439f9.pth)	hf_hub_idr  zregnety_040.ra3_in1kzshttps://github.com/huggingface/pytorch-image-models/releases/download/v0.1-tpu-weights/regnety_040_ra3-670e1166.pthzregnety_064.ra3_in1kzshttps://github.com/huggingface/pytorch-image-models/releases/download/v0.1-tpu-weights/regnety_064_ra3-aa26dc7d.pthzregnety_080.ra3_in1kzshttps://github.com/huggingface/pytorch-image-models/releases/download/v0.1-tpu-weights/regnety_080_ra3-1fdc4344.pthzregnety_120.sw_in12k_ft_in1k)r  zregnety_160.sw_in12k_ft_in1kzregnety_160.lion_in12k_ft_in1kzregnety_120.sw_in12ki-.  )r  r   zregnety_160.sw_in12kzregnety_040_sgn.untrained)r  zregnetv_040.ra3_in1kzshttps://github.com/huggingface/pytorch-image-models/releases/download/v0.1-tpu-weights/regnetv_040_ra3-c248f51f.pthr   )r  r  r  zregnetv_064.ra3_in1kzshttps://github.com/huggingface/pytorch-image-models/releases/download/v0.1-tpu-weights/regnetv_064_ra3-530616c2.pthzregnetz_005.untrainedzregnetz_040.ra3_in1kzshttps://github.com/huggingface/pytorch-image-models/releases/download/v0.1-tpu-weights/regnetz_040_ra3-9007edf5.pth)r      r  )r`   r`   r1   )r   r  r  )r  r  r  r  r  r  zregnetz_040_h.ra3_in1kzthttps://github.com/huggingface/pytorch-image-models/releases/download/v0.1-tpu-weights/regnetz_040h_ra3-f594343b.pthzregnety_160.deit_in1kz<https://dl.fbaipublicfiles.com/deit/regnety_160-a5fe301d.pthzregnetx_004_tv.tv2_in1kz?https://download.pytorch.org/models/regnet_x_400mf-62229a5f.pthzregnetx_008.tv2_in1kz?https://download.pytorch.org/models/regnet_x_800mf-94a99ebd.pthzregnetx_016.tv2_in1kz?https://download.pytorch.org/models/regnet_x_1_6gf-a12f2b72.pthzregnetx_032.tv2_in1kz?https://download.pytorch.org/models/regnet_x_3_2gf-7071aa85.pthzregnetx_080.tv2_in1kz=https://download.pytorch.org/models/regnet_x_8gf-2b70d774.pthzregnetx_160.tv2_in1kz>https://download.pytorch.org/models/regnet_x_16gf-ba3796d7.pthzregnetx_320.tv2_in1kz>https://download.pytorch.org/models/regnet_x_32gf-6eb8fdc6.pthzregnety_004.tv2_in1kz?https://download.pytorch.org/models/regnet_y_400mf-e6988f5f.pthzregnety_008_tv.tv2_in1kz?https://download.pytorch.org/models/regnet_y_800mf-58fc7688.pthzregnety_016.tv2_in1kz?https://download.pytorch.org/models/regnet_y_1_6gf-0d7bc02a.pthzregnety_032.tv2_in1kz?https://download.pytorch.org/models/regnet_y_3_2gf-9180c971.pthzregnety_080_tv.tv2_in1kz=https://download.pytorch.org/models/regnet_y_8gf-dc2b1b54.pthzregnety_160.tv2_in1kz>https://download.pytorch.org/models/regnet_y_16gf-3e4a00f9.pthzregnety_320.tv2_in1kz>https://download.pytorch.org/models/regnet_y_32gf-8db6d4b5.pthzregnety_160.swag_ft_in1kzChttps://download.pytorch.org/models/regnet_y_16gf_swag-43afe44d.pthzcc-by-nc-4.0)r     r  )   r  )r  r  r  r  r  r  zregnety_320.swag_ft_in1kzChttps://download.pytorch.org/models/regnet_y_32gf_swag-04fdfa75.pthzregnety_1280.swag_ft_in1kzDhttps://download.pytorch.org/models/regnet_y_128gf_swag-c8ce3e52.pthzregnety_160.swag_lc_in1kzFhttps://download.pytorch.org/models/regnet_y_16gf_lc_swag-f3ec0043.pth)r  r  r  zregnety_320.swag_lc_in1kzFhttps://download.pytorch.org/models/regnet_y_32gf_lc_swag-e1583746.pthzregnety_1280.swag_lc_in1kzGhttps://download.pytorch.org/models/regnet_y_128gf_lc_swag-cbe8ce12.pthzregnety_320.seer_ft_in1kzseer-licensez)https://github.com/facebookresearch/visslzhttps://dl.fbaipublicfiles.com/vissl/model_zoo/seer_finetuned/seer_regnet32_finetuned_in1k_model_final_checkpoint_phase78.torch)r  r  r  r  r  r  r  zregnety_640.seer_ft_in1kzhttps://dl.fbaipublicfiles.com/vissl/model_zoo/seer_finetuned/seer_regnet64_finetuned_in1k_model_final_checkpoint_phase78.torchzregnety_1280.seer_ft_in1kzhttps://dl.fbaipublicfiles.com/vissl/model_zoo/seer_finetuned/seer_regnet128_finetuned_in1k_model_final_checkpoint_phase78.torchzregnety_2560.seer_ft_in1kzhttps://dl.fbaipublicfiles.com/vissl/model_zoo/seer_finetuned/seer_regnet256_finetuned_in1k_model_final_checkpoint_phase38.torchzregnety_320.seerzihttps://dl.fbaipublicfiles.com/vissl/model_zoo/seer_regnet32d/seer_regnet32gf_model_iteration244000.torch)r  r  r   r  r  zregnety_640.seerzphttps://dl.fbaipublicfiles.com/vissl/model_zoo/seer_regnet64/seer_regnet64gf_model_final_checkpoint_phase0.torchzregnety_1280.seerzhttps://dl.fbaipublicfiles.com/vissl/model_zoo/swav_ig1b_regnet128Gf_cnstant_bs32_node16_sinkhorn10_proto16k_syncBN64_warmup8k/model_final_checkpoint_phase0.torchzregnetx_002.pycls_in1kzregnetx_004.pycls_in1kzregnetx_006.pycls_in1kzregnetx_008.pycls_in1kzregnetx_016.pycls_in1kzregnetx_032.pycls_in1kzregnetx_040.pycls_in1kzregnetx_064.pycls_in1kzregnetx_080.pycls_in1kzregnetx_120.pycls_in1kzregnetx_160.pycls_in1kzregnetx_320.pycls_in1kzregnety_002.pycls_in1kzregnety_004.pycls_in1kzregnety_006.pycls_in1kzregnety_008.pycls_in1kzregnety_016.pycls_in1kzregnety_032.pycls_in1kzregnety_040.pycls_in1kzregnety_064.pycls_in1kzregnety_080.pycls_in1kzregnety_120.pycls_in1kzregnety_160.pycls_in1kzregnety_320.pycls_in1kc                     t        d| fi |S )zRegNetX-200MFry  r  r  r   s     rL   ry  ry         ->v>>rK   c                     t        d| fi |S )zRegNetX-400MFr{  r  r  s     rL   r{  r{    r  rK   c                     t        d| fi |S )z+RegNetX-400MF w/ torchvision group roundingr}  r  r  s     rL   r}  r}         *JA&AArK   c                     t        d| fi |S )zRegNetX-600MFr~  r  r  s     rL   r~  r~    r  rK   c                     t        d| fi |S )zRegNetX-800MFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetX-1.6GFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetX-3.2GFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetX-4.0GFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetX-6.4GFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetX-8.0GFr  r  r  s     rL   r  r  %  r  rK   c                     t        d| fi |S )zRegNetX-12GFr  r  r  s     rL   r  r  +  r  rK   c                     t        d| fi |S )zRegNetX-16GFr  r  r  s     rL   r  r  1  r  rK   c                     t        d| fi |S )zRegNetX-32GFr  r  r  s     rL   r  r  7  r  rK   c                     t        d| fi |S )zRegNetY-200MFr  r  r  s     rL   r  r  =  r  rK   c                     t        d| fi |S )zRegNetY-400MFr  r  r  s     rL   r  r  C  r  rK   c                     t        d| fi |S )zRegNetY-600MFr  r  r  s     rL   r  r  I  r  rK   c                     t        d| fi |S )zRegNetY-800MFr  r  r  s     rL   r  r  O  r  rK   c                     t        d| fi |S )z+RegNetY-800MF w/ torchvision group roundingr  r  r  s     rL   r  r  U  r  rK   c                     t        d| fi |S )zRegNetY-1.6GFr  r  r  s     rL   r  r  [  r  rK   c                     t        d| fi |S )zRegNetY-3.2GFr  r  r  s     rL   r  r  a  r  rK   c                     t        d| fi |S )zRegNetY-4.0GFr  r  r  s     rL   r  r  g  r  rK   c                     t        d| fi |S )zRegNetY-6.4GFr  r  r  s     rL   r  r  m  r  rK   c                     t        d| fi |S )zRegNetY-8.0GFr  r  r  s     rL   r  r  s  r  rK   c                     t        d| fi |S )z+RegNetY-8.0GF w/ torchvision group roundingr  r  r  s     rL   r  r  y  r  rK   c                     t        d| fi |S )zRegNetY-12GFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetY-16GFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetY-32GFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetY-64GFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetY-128GFr  r  r  s     rL   r  r         .*???rK   c                     t        d| fi |S )zRegNetY-256GFr  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetY-4.0GF w/ GroupNorm r  r  r  s     rL   r  r    s     +ZB6BBrK   c                     t        d| fi |S )zRegNetV-4.0GF (pre-activation)r  r  r  s     rL   r  r    r  rK   c                     t        d| fi |S )zRegNetV-6.4GF (pre-activation)r  r  r  s     rL   r  r    r  rK   c                      t        d| fddi|S )zRegNetZ-500MF
    NOTE: config found in https://github.com/facebookresearch/ClassyVision/blob/main/classy_vision/models/regnet.py
    but it's not clear it is equivalent to paper model as not detailed in the paper.
    r  r   Fr  r  s     rL   r  r         -TETVTTrK   c                      t        d| fddi|S )RegNetZ-4.0GF
    NOTE: config found in https://github.com/facebookresearch/ClassyVision/blob/main/classy_vision/models/regnet.py
    but it's not clear it is equivalent to paper model as not detailed in the paper.
    r  r   Fr  r  s     rL   r  r    r  rK   c                      t        d| fddi|S )r	  r  r   Fr  r  s     rL   r  r    s     /:VeVvVVrK   regnetz_040h)r3   )r`   )r   r   r   NFNN)r   NFNN)r   FrJ   )r   rG  )mrD   rk   dataclassesr   r   	functoolsr   typingr   r   r   r	   r
   r   r   r   rg   torch.nnr   	timm.datar   r   timm.layersr   r   r   r   r   r   r   r   r   r   r   _builderr   	_featuresr   _manipulater   r    	_registryr!   r"   r#   __all__r%   rG   rE   rR   r_   rz   r   rI   r   r   r   rH   r   r   r   r   r$   r   rx  r   r  r  r  r  r  default_cfgsry  r{  r}  r~  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  rA   rJ   rK   rL   <module>r     s  2  *  J J J   A    X X * + 4 Y Y[
! 3 3 3&
!e 
! 
! 
!" 	S	E{ S	 	
 49d3i D >>> > 	>
 > > 49c49$%>H 040
0
0
 0
 	0

 0
 T"))_-0
 0
 YY0
l 04&(&(&( &( 	&(
 &( T"))_-&( &( ]]&(^ %+04&!#&& & 	&
 & S/& T"))_-& & bii&Rg gTgBII gTBryy BJhRYY hV	 "))  3  T  VZ  ,>4S> >d38n >D  =REdqK= REdrL= u"B`cd	=
 REdrL= REdrL= REdrL= REdrL= REdrL= SUt"M= REds"M= SUt2N= SUss"M= SUss"M=" REdqUYZ#=$ REdqUYZ%=& REdrVZ['=( REcbUYZ)=* u2X\nqr+=, REdrVZ[-=. REdrVZ[/=0 REdrVZ[1=2 SUt"W[\3=4 SUt"W[\5=6 $2RZ^pst7=8 SUt2X\]9=: SVBY]^;=< SVBY]^==> SV2X\]?=@ cf#RZ^_A=B cf#RZ^_C=J %DRrDW\b%IKK=T REdrDQUagiU=X SUtTRVbhY=b RDTacTXDtvc=j RDUqsUYDqFk=r RDUqsUYDtvs=
@C T  $c T#s(^ & S#X & S#X & % u&4|~u&
 D BCu& D BCu& D BCu& #D7$;u& #D7$;u&  %dW&=!u&& D'u&, D-u&6  "7u&8 D B9u&@ D BAu&J Tb\Ku&L D B FSR_aMu&T d C FSR_aUu&` T]_au&f wM Ogu&l GMOmu&r GMOsu&x GMOyu&~ GKMu&D GLNEu&J GLNKu&R GMOSu&X wM OYu&^ GMO_u&d GMOeu&j wK Mku&p GLNqu&v GLNwu&~ Q[i Hs!Du&F Q[i Hs!DGu&N  R\j Hs"DOu&X T^l!nYu&^ T^l!n_u&d  U_m"oeu&l +V N Hs	!Dmu&v +V N Hs	!Dwu&@  +V O Hs	"DAu&J  +V O Hs	"DKu&V w~:egWu&^ ~~:eg_u&f  q~:eggu&x g8yu&z g8{u&| g8}u&~ g8u&@ g8Au&B g8Cu&D g8Eu&F g8Gu&H g8Iu&J g8Ku&L g8Mu&N g8Ou&R g8Su&T g8Uu&V g8Wu&X g8Yu&Z g8[u&\ g8]u&^ g8_u&` g8au&b g8cu&d g8eu&f g8gu&h g8iu& up ?D ?v ? ?
 ?D ?v ? ?
 Bt B& B B
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 Bt B& B B
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 Bt B& B B
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 ?D ?v ? ?
 @T @ @ @
 @T @ @ @
 C C6 C C
 ?D ?v ? ?
 ?D ?v ? ?
 UD Uv U U UD Uv U U Wd W W W HO' rK   