
    ^j                        d Z ddlmZ ddlmZ ddl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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'm(Z(m)Z) ddl*m+Z+m,Z,m-Z- dgZ. G d dej^                        Z0 G d dej^                        Z1 G d dej^                        Z2 G d dej^                        Z3 G d dej^                        Z4 G d dej^                        Z5de6de7fdZ8ddde eed !      ddfd"e9d#e9de6d$e7d%ed&edejt                  fd'Z; G d( dej^                        Z<dwd)ej^                  d*e6d+e7ddfd,Z= ej|                         dxd-ej^                  d.e6d/e6fd0       Z?dyd1e6d2e7d3ede<fd4Z@dyd1e6d2e7d3ede<fd5ZAdzd6e6d3edee6ef   fd7ZB e+i d8 eBd9d:d;      d< eBd9d:d;      d= eBd9d>d?d@d:dA      dB eBd9dCdDd@dE      dF eBd9dCdDd@dE      dG eBd9dCdDd@dE      dH eBd9dCdDd@dE      dI eBd9dCdDd@dE      dJ eBd9dKdLd@dE      dM eBd9dNdO      dP eBd9dNdO      dQ eBd9dNdO      dR eBd9dNdO      dS eBd9dNdO      dT eBd9dNdO      dU eBd9d:dVdWd@X      dY eBd9d:dVdWd@dZ[       eBd9d:dVdWd@X       eBd9d:dVdWd@dZ[       eBd9d@d>d?dCd:dZ\       eBd9d:d]dWd@X       eBd:dZ^       eBd:dZ^       eBd9d:d]dWd@X       eBd:dZ^       eBd:_       eBd:dZ^       eBd9d:dZd]dWd@`       eBd9d:dZd]dWd@`       eBd:dZ^      da      ZCe,dyd2e7d3ede<fdb       ZDe,dyd2e7d3ede<fdc       ZEe,dyd2e7d3ede<fdd       ZFe,dyd2e7d3ede<fde       ZGe,dyd2e7d3ede<fdf       ZHe,dyd2e7d3ede<fdg       ZIe,dyd2e7d3ede<fdh       ZJe,dyd2e7d3ede<fdi       ZKe,dyd2e7d3ede<fdj       ZLe,dyd2e7d3ede<fdk       ZMe,dyd2e7d3ede<fdl       ZNe,dyd2e7d3ede<fdm       ZOe,dyd2e7d3ede<fdn       ZPe,dyd2e7d3ede<fdo       ZQe,dyd2e7d3ede<fdp       ZRe,dyd2e7d3ede<fdq       ZSe,dyd2e7d3ede<fdr       ZTe,dyd2e7d3ede<fds       ZUe,dyd2e7d3ede<fdt       ZVe,dyd2e7d3ede<fdu       ZW e-eXdBdFdGdHdIdJdMdPdQdRdSdTd8d<d=dv       y){a/  Pre-Activation ResNet v2 with GroupNorm and Weight Standardization.

A PyTorch implementation of ResNetV2 adapted from the Google Big-Transfer (BiT) source code
at https://github.com/google-research/big_transfer to match timm interfaces. The BiT weights have
been included here as pretrained models from their original .NPZ checkpoints.

Additionally, supports non pre-activation bottleneck for use as a backbone for Vision Transformers (ViT) and
extra padding support to allow porting of official Hybrid ResNet pretrained weights from
https://github.com/google-research/vision_transformer

Thanks to the Google team for the above two repositories and associated papers:
* Big Transfer (BiT): General Visual Representation Learning - https://arxiv.org/abs/1912.11370
* An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale - https://arxiv.org/abs/2010.11929
* Knowledge distillation: A good teacher is patient and consistent - https://arxiv.org/abs/2106.05237

Original copyright of Google code below, modifications by Ross Wightman, Copyright 2020.
    )OrderedDict)partial)AnyCallableDictListOptionalTupleUnionNIMAGENET_INCEPTION_MEANIMAGENET_INCEPTION_STD)GroupNormActBatchNormAct2dEvoNorm2dS0FilterResponseNormTlu2dClassifierHeadDropPathcalculate_drop_path_ratesAvgPool2dSamecreate_pool2d	StdConv2dcreate_conv2dget_act_layerget_norm_act_layermake_divisible   )build_model_with_cfg)feature_take_indices)checkpoint_seqnamed_applyadapt_input_conv)generate_default_cfgsregister_modelregister_model_deprecationsResNetV2c                        e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 ddedee   dedededee   ded	ee   d
ee   dee   dee   def fdZddZ	de
j                  de
j                  fdZ xZS )PreActBasiczAPre-activation basic block (not in typical 'v2' implementations).in_chsout_chsbottle_ratiostridedilationfirst_dilationgroups	act_layer
conv_layer
norm_layer
proj_layerdrop_path_ratec           
         ||d}t         |           |xs |}|	xs t        }	|
xs t        t        d      }
|xs |}t        ||z        }|&|dk7  s
||k7  s||k7  r |||f|||d|	|
d|| _        nd| _         |
|fi || _         |	||df|||d	|| _         |
|fi || _	         |	||df||d
|| _
        |dkD  rt        |      | _        yt        j                         | _        y)aw  Initialize PreActBasic block.

        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            bottle_ratio: Bottleneck ratio (not used in basic block).
            stride: Stride for convolution.
            dilation: Dilation rate.
            first_dilation: First dilation rate.
            groups: Group convolution size.
            act_layer: Activation layer type.
            conv_layer: Convolution layer type.
            norm_layer: Normalization layer type.
            proj_layer: Projection/downsampling layer type.
            drop_path_rate: Stochastic depth drop rate.
        devicedtype    
num_groupsNr   Tr,   r-   r.   preactr1   r2      r,   r-   r/   )r-   r/   r   )super__init__r   r   r   r   
downsamplenorm1conv1norm2conv2r   nnIdentity	drop_pathselfr)   r*   r+   r,   r-   r.   r/   r0   r1   r2   r3   r4   r7   r8   ddmid_chs	__class__s                    _/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/resnetv2.pyrA   zPreActBasic.__init__5   s1   B /'38,9
G7<B#G
#V <!78!v{n6PTZ^eTe(
 !-%%
 
DO #DO-"-
p6Ncipmop
.2.
!\hv\Y[\
5Ca5G.1R[[]    returnc                 j    t         j                  j                  | j                  j                         y)zLZero-initialize the last convolution weight (not applicable to basic block).N)rG   initzeros_rF   weightrK   s    rO   zero_init_lastzPreActBasic.zero_init_lasts       
tzz(()rP   xc                     | j                  |      }|}| j                  | j                  |      }| j                  |      }| j                  | j	                  |            }| j                  |      }||z   S zoForward pass.

        Args:
            x: Input tensor.

        Returns:
            Output tensor.
        )rC   rB   rD   rF   rE   rI   rK   rY   x_preactshortcuts       rO   forwardzPreActBasic.forwardw   sn     ::a= ??&x0H JJx JJtzz!}%NN18|rP   )N      ?r   r   Nr   NNNN        NNrQ   N__name__
__module____qualname____doc__intr	   floatr   rA   rW   torchTensorr_   __classcell__rN   s   @rO   r(   r(   2   s    K
 &*"%,0,0-1-1-1$&<[<[ c]<[  	<[
 <[ <[ %SM<[ <[  )<[ !*<[ !*<[ !*<[ "<[|* %,, rP   r(   c                        e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 ddedee   dedededee   ded	ee   d
ee   dee   dee   def fdZddZ	de
j                  de
j                  fdZ xZS )PreActBottlenecka  Pre-activation (v2) bottleneck block.

    Follows the implementation of "Identity Mappings in Deep Residual Networks":
    https://github.com/KaimingHe/resnet-1k-layers/blob/master/resnet-pre-act.lua

    Except it puts the stride on 3x3 conv when available.
    r)   r*   r+   r,   r-   r.   r/   r0   r1   r2   r3   r4   c           
         ||d}t         |           |xs |}|	xs t        }	|
xs t        t        d      }
|xs |}t        ||z        }| |||f|||d|	|
d|| _        nd| _         |
|fi || _         |	||dfi || _         |
|fi || _	         |	||df|||d	|| _
         |
|fi || _         |	||dfi || _        |d
kD  rt        |      | _        yt        j                         | _        y)ab  Initialize PreActBottleneck block.

        Args:
            in_chs: Input channels.
            out_chs: Output channels.
            bottle_ratio: Bottleneck ratio.
            stride: Stride for convolution.
            dilation: Dilation rate.
            first_dilation: First dilation rate.
            groups: Group convolution size.
            act_layer: Activation layer type.
            conv_layer: Convolution layer type.
            norm_layer: Normalization layer type.
            proj_layer: Projection/downsampling layer type.
            drop_path_rate: Stochastic depth drop rate.
        r6   r9   r:   NTr<   r   r>   r?   r   )r@   rA   r   r   r   r   rB   rC   rD   rE   rF   norm3conv3r   rG   rH   rI   rJ   s                    rO   rA   zPreActBottleneck.__init__   s9   B /'38,9
G7<B#G
#V <!78!(
 !-%%
 
DO #DO-"-
9b9
.2.
!qF^djqnpq
.2.
!:r:
5Ca5G.1R[[]rP   rQ   c                 j    t         j                  j                  | j                  j                         y)z,Zero-initialize the last convolution weight.N)rG   rS   rT   rr   rU   rV   s    rO   rW   zPreActBottleneck.zero_init_last   rX   rP   rY   c                 0   | j                  |      }|}| j                  | j                  |      }| j                  |      }| j                  | j	                  |            }| j                  | j                  |            }| j                  |      }||z   S r[   )rC   rB   rD   rF   rE   rr   rq   rI   r\   s       rO   r_   zPreActBottleneck.forward   s     ::a= ??&x0H JJx JJtzz!}%JJtzz!}%NN18|rP   N      ?r   r   Nr   NNNNra   NNrb   rc   rm   s   @rO   ro   ro      s     &*"&,0,0-1-1-1$&>[>[ c]>[  	>[
 >[ >[ %SM>[ >[  )>[ !*>[ !*>[ !*>[ ">[@* %,, rP   ro   c                        e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 ddedee   dedededee   ded	ee   d
ee   dee   dee   def fdZddZ	de
j                  de
j                  fdZ xZS )
BottleneckzUNon Pre-activation bottleneck block, equiv to V1.5/V1b Bottleneck. Used for ViT.
    r)   r*   r+   r,   r-   r.   r/   r0   r1   r2   r3   r4   c           	      2   ||d}t         |           |xs |}|xs t        j                  }|	xs t        }	|
xs t        t        d      }
|xs |}t        ||z        }| |||f||d|	|
d|| _        nd | _         |	||dfi || _	         |
|fi || _
         |	||df|||d|| _         |
|fi || _         |	||dfi || _         |
|fd	di|| _        |d
kD  rt        |      nt        j                          | _         |d      | _        y )Nr6   r9   r:   F)r,   r-   r=   r1   r2   r   r>   r?   	apply_actr   T)inplace)r@   rA   rG   ReLUr   r   r   r   rB   rD   rC   rF   rE   rr   rq   r   rH   rI   act3rJ   s                    rO   rA   zBottleneck.__init__   sP   " /'38(	,9
G7<B#G
#V <!78!(	 !%%	 	DO #DO9b9
.2.
!qF^djqnpq
.2.
!:r:
?5?B?
5Ca5G.1R[[]d+	rP   rQ   c                     t        | j                  dd      4t        j                  j	                  | j                  j
                         yy)z+Zero-initialize the last batch norm weight.rU   N)getattrrq   rG   rS   rT   rU   rV   s    rO   rW   zBottleneck.zero_init_last'  s4    4::x.:GGNN4::,,- ;rP   rY   c                 Z   |}| j                   | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }| j                  |      }| j                  |      }| j                  ||z         }|S r[   )	rB   rD   rC   rF   rE   rr   rq   rI   r}   )rK   rY   r^   s      rO   r_   zBottleneck.forward,  s     ??&q)H JJqMJJqMJJqMJJqMJJqMJJqMNN1IIa(l#rP   ru   rb   rc   rm   s   @rO   rx   rx      s    
 &*"&,0,0-1-1-1$&/,/, c]/,  	/,
 /, /, %SM/, /,  )/, !*/, !*/, !*/, "/,b.
 %,, rP   rx   c                        e Zd ZdZ	 	 	 	 	 	 	 	 ddededededee   dedee   d	ee   f fd
Zde	j                  de	j                  fdZ xZS )DownsampleConvz$1x1 convolution downsampling module.r)   r*   r,   r-   r.   r=   r1   r2   c                     |	|
d}t         |            |||dfd|i|| _        |rt        j                         | _        y  ||fddi|| _        y )Nr6   r   r,   rz   F)r@   rA   convrG   rH   norm)rK   r)   r*   r,   r-   r.   r=   r1   r2   r7   r8   rL   rN   s               rO   rA   zDownsampleConv.__init__I  sY     /vwG&GBG	%+BKKM	G1[u1[XZ1[	rP   rY   rQ   c                 B    | j                  | j                  |            S ztForward pass.

        Args:
            x: Input tensor.

        Returns:
            Downsampled tensor.
        )r   r   rK   rY   s     rO   r_   zDownsampleConv.forward[  s     yy1&&rP   r   r   NTNNNNrd   re   rf   rg   rh   r	   boolr   rA   rj   rk   r_   rl   rm   s   @rO   r   r   F  s    . ,0-1-1\\ \ 	\
 \ %SM\ \ !*\ !*\$	' 	'%,, 	'rP   r   c                        e Zd ZdZ	 	 	 	 	 	 	 	 ddededededee   dedee   d	ee   f fd
Zde	j                  de	j                  fdZ xZS )DownsampleAvgz/AvgPool downsampling as in 'D' ResNet variants.r)   r*   r,   r-   r.   r=   r1   r2   c                 d   |	|
d}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fddi|| _        |rt        j                         | _        y  ||fddi|| _        y )	Nr6   r      TF)	ceil_modecount_include_padr,   rz   )	r@   rA   r   rG   	AvgPool2dpoolrH   r   r   )rK   r)   r*   r,   r-   r.   r=   r1   r2   r7   r8   rL   
avg_strideavg_pool_fnrN   s                 rO   rA   zDownsampleAvg.__init__j  s     /'1}V!
A:A+5?x!|-QSQ]Q]K#AzTUZ[DIDIvwB!BrB	%+BKKM	G1[u1[XZ1[	rP   rY   rQ   c                 `    | j                  | j                  | j                  |                  S r   )r   r   r   r   s     rO   r_   zDownsampleAvg.forward  s$     yy499Q<011rP   r   r   rm   s   @rO   r   r   g  s    9 ,0-1-1\\ \ 	\
 \ %SM\ \ !*\ !*\0	2 	2%,, 	2rP   r   c                        e Zd ZdZddddedddfdededed	ed
ededededee	e      de
dee
   dee
   dee
   def fdZdej                  dej                  fdZ xZS )ResNetStagezResNet Stage.rv   r   FNr)   r*   r,   r-   depthr+   r/   avg_down	block_dprblock_fnr0   r1   r2   block_kwargsc                 l   t         |           d| _        |dv rdnd}t        |||      }|rt        nt
        }|}t        j                         | _        t        |      D ]Q  }|	r|	|   nd}|dk(  r|nd}| j                  j                  t        |       |
||f|||||||d||       |}|}d }S y )	NF)r   r   r   r   )r0   r1   r2   ra   r   )r,   r-   r+   r/   r.   r3   r4   )r@   rA   grad_checkpointingdictr   r   rG   
Sequentialblocksrange
add_modulestr)rK   r)   r*   r,   r-   r   r+   r/   r   r   r   r0   r1   r2   r   r.   layer_kwargsr3   prev_chs	block_idxr4   rN   s                        rO   rA   zResNetStage.__init__  s    " 	"'&&0aiJS]^&.]N
mmou 	I5>Yy1BN(A~V1FKK""3y>84 !)-%-4 4 4  H%NJ%	rP   rY   rQ   c                     | j                   r6t        j                  j                         st	        | j
                  |      }|S | j                  |      }|S )zForward pass through all blocks in the stage.

        Args:
            x: Input tensor.

        Returns:
            Output tensor.
        )r   rj   jitis_scriptingr    r   r   s     rO   r_   zResNetStage.forward  sG     ""599+A+A+Ct{{A.A  AArP   )rd   re   rf   rg   ro   rh   ri   r   r	   r   r   r   rA   rj   rk   r_   rl   rm   s   @rO   r   r     s     #'"/3!1,0-1-1++ + 	+
 + +  + + +  U,+ +  )+ !*+ !*+  +Z %,, rP   r   	stem_typerQ   c                 B    t        dD cg c]  }|| v  c}      S c c}w )zCheck if stem type is deep (has multiple convolutions).

    Args:
        stem_type: Type of stem to check.

    Returns:
        True if stem is deep, False otherwise.
    )deeptiered)any)r   ss     rO   is_stem_deepr     s"     (:;1Y;<<;s   @    Tr9   r:   r)   r*   r=   r1   r2   c                    ||d}t               }	|dv sJ t        |      rd|v rd|z  dz  |dz  f}
n
|dz  |dz  f}
 || |
d   fddd||	d	<    ||
d   fi ||	d
<    ||
d   |
d   fddd||	d<    ||
d   fi ||	d<    ||
d   |fddd||	d<   |s+ ||fi ||	d<   n || |fddd||	d<   |s ||fi ||	d<   d|v r5t        j                  dd      |	d<   t        j                  ddd      |	d<   n2d|v rt        dddd      |	d<   nt        j                  ddd      |	d<   t        j                  |	      S )Nr6   )r   fixedsamer   
deep_fixed	deep_samer   r   r>      r   r   )kernel_sizer,   rD   rC   r   rF   rE   rr   rq      r   r   r   ra   pad)r   r,   paddingr   r   max)r   r   rG   ConstantPad2d	MaxPool2dr   r   )r)   r*   r   r=   r1   r2   r7   r8   rL   stemstem_chss              rO   create_resnetv2_stemr     s    U	+B=DZZZZ Iy Gq('Q,7H1gl3H"68A;VAaVSUVW"8A;5"5W"8A;[ST[XZ[W"8A;5"5W"8A;WQqWTVWW&w5"5DM "&'QqQbQV%g44DL)&&q"-U||!QGV	9	$U!VTV ||!QGV==rP   c            '           e Zd ZdZdddddddd	d
dd
dej
                   eed      eddd
ddfde	e
   dee
df   de
de
dede
de
de
dededededededed ed!ed"ed#ef& fd$Zej$                  j&                  d>d#ed%dfd&       Zej$                  j'                         d?d'ed(ed%d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j4                  fd.       ZdAde
dee   d%dfd/Z	 	 	 	 	 dBd0ej<                  d1eee
e	e
   f      d2ed3ed4ed5ed%ee	ej<                     eej<                  e	ej<                     f   f   fd6Z 	 	 	 dCd1ee
e	e
   f   d7ed8efd9Z!d0ej<                  d%ej<                  fd:Z"d@d0ej<                  d;ed%ej<                  fd<Z#d0ej<                  d%ej<                  fd=Z$ xZ%S )Dr&   z7Implementation of Pre-activation (v2) ResNet mode.
    )      i   i     r>   avgr9   r   r   r   FTrv   r:   ra   Nlayerschannels.num_classesin_chansglobal_pooloutput_stridewidth_factorr   r   r   r=   basicr+   r0   r2   r1   	drop_rater4   rW   c                 2   t         %|           ||d}|| _        || _        || _        |}t        ||      }t        |      }g | _        t        ||z        }t        |||	|f||d|| _
        |rt        |	      rdndnd}| j                  j                  t        |d|             |}d	}d
}t        ||d      }|r|rt        nt         }n
|rJ t"        }t%        j&                         | _        t+        t-        |||            D ]  \  }\  }} }!t        | |z        }"|dk(  rd
nd}#||k\  r||#z  }d
}#t/        ||"f|#||||
||||!|d
|}$|"}||#z  }| xj                  t        ||d|       gz  c_        | j(                  j1                  t3        |      |$        |x| _        | _        |r || j4                  fi |nt%        j8                         | _        t=        | j4                  |f|| j                  dd|| _        | jA                  |       y)a  
        Args:
            layers (List[int]) : number of layers in each block
            channels (List[int]) : number of channels in each block:
            num_classes (int): number of classification classes (default 1000)
            in_chans (int): number of input (color) channels. (default 3)
            global_pool (str): Global pooling type. One of 'avg', 'max', 'avgmax', 'catavgmax' (default 'avg')
            output_stride (int): output stride of the network, 32, 16, or 8. (default 32)
            width_factor (int): channel (width) multiplication factor
            stem_chs (int): stem width (default: 64)
            stem_type (str): stem type (default: '' == 7x7)
            avg_down (bool): average pooling in residual downsampling (default: False)
            preact (bool): pre-activation (default: True)
            act_layer (Union[str, nn.Module]): activation layer
            norm_layer (Union[str, nn.Module]): normalization layer
            conv_layer (nn.Module): convolution module
            drop_rate: classifier dropout rate (default: 0.)
            drop_path_rate: stochastic depth rate (default: 0.)
            zero_init_last: zero-init last weight in residual path (default: False)
        r6   )r0   )r1   r2   z
stem.conv3	stem.convz	stem.normr   )num_chs	reductionmodule   r   T)	stagewiser   )
r,   r-   r   r+   r   r0   r1   r2   r   r   zstages.)	pool_typer   use_convrW   N)!r@   rA   r   r   r   r   r   feature_infor   r   r   r   appendr   r   r(   ro   rx   rG   r   stages	enumeratezipr   r   r   num_featureshead_hidden_sizerH   r   r   headinit_weights)&rK   r   r   r   r   r   r   r   r   r   r   r=   r   r+   r0   r2   r1   r   r4   rW   r7   r8   rL   wf	stem_featr   curr_strider-   
block_dprsr   	stage_idxdcbdprr*   r,   stagerN   s&                                        rO   rA   zResNetV2.__init__  so   X 	/& "'
iH
!),	!(R-0(	

 "!
 
	 SY\)%<\+^i	  h!I!VW.~vQUV
&+{1AH9!Hmmo'0VXz1R'S 	:#I|1d$QV,G#q.QaFm+F" !)!#%%! E H6!K$x;Y`aj`kWl"m!nnKK""3y>591	:4 5=<D1;AJt007B7r{{}	"
 "nn
 
	 	8rP   rQ   c                 :    t        t        t        |      |        y)zInitialize model weights.r   N)r!   r   _init_weights)rK   rW   s     rO   r   zResNetV2.init_weights  s     	GM.I4PrP   checkpoint_pathprefixc                     t        | ||       y)zLoad pretrained weights.N)_load_weights)rK   r   r   s      rO   load_pretrainedzResNetV2.load_pretrained  s     	dOV4rP   coarsec                 ,    t        d|rdnddg      }|S )z"Group parameters for optimization.z^stemz^stages\.(\d+))z^stages\.(\d+)\.blocks\.(\d+)N)z^norm)i )r   r   )r   )rK   r   matchers      rO   group_matcherzResNetV2.group_matcher  s)     (.$8$5
 rP   enablec                 4    | j                   D ]	  }||_         y)z)Enable or disable gradient checkpointing.N)r   r   )rK   r   r   s      rO   set_grad_checkpointingzResNetV2.set_grad_checkpointing  s      	*A#)A 	*rP   c                 .    | j                   j                  S )zGet the classifier head.)r   fcrV   s    rO   get_classifierzResNetV2.get_classifier  s     yy||rP   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)rK   r   r   s      rO   reset_classifierzResNetV2.reset_classifier  s     '		[1rP   rY   indicesr   
stop_early
output_fmtintermediates_onlyc                    |dv sJ d       g }t        d|      \  }}	d}
|j                  dd \  }}| j                  D ]'  } ||      }|j                  dd |dz  |dz  fk(  s&|}) |
|v r|j                         t	        | j
                        }t        j                  j                         s|s| j
                  }n| j
                  d|	 }t        |d	      D ]O  \  }
} ||      }|
|v s|
|k(  r'|r| j                  |      n|}|j                  |       ?|j                  |       Q |r|S |
|k(  r| j                  |      }||fS )
a   Forward features that returns intermediates.

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

        )NCHWzOutput shape must be NCHW.   r   Nr   r   )start)r   shaper   r   lenr   rj   r   r   r   r   )rK   rY   r  r   r  r  r  intermediatestake_indices	max_indexfeat_idxHWr   x_downlast_idxr   r   x_inters                      rO   forward_intermediateszResNetV2.forward_intermediates  sc   * Y&D(DD&"6q'"Bi wwrs|1II 	DQAwwrs|Aq1u~-	 |#  (t{{#99!!#:[[F[[),F(q9 	,OHeaA<'x'.2diilG!((1!((+	,   x		!A-rP   
prune_norm
prune_headc                     t        d|      \  }}| j                  d| | _        |rt        j                         | _        |r| j                  dd       |S )z@ Prune layers not required for specified intermediates.
        r	  Nr   r   )r   r   rG   rH   r   r  )rK   r  r  r  r  r  s         rO   prune_intermediate_layersz"ResNetV2.prune_intermediate_layers  sP     #7q'"Bikk*9-DI!!!R(rP   c                 l    | j                  |      }| j                  |      }| j                  |      }|S )zForward pass through feature extraction layers.

        Args:
            x: Input tensor.

        Returns:
            Feature tensor.
        )r   r   r   r   s     rO   forward_featureszResNetV2.forward_features  s0     IIaLKKNIIaLrP   
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   )rK   rY   r  s      rO   forward_headzResNetV2.forward_head  s(     7Atyyzy2RdiiPQlRrP   c                 J    | j                  |      }| j                  |      }|S )zoForward pass.

        Args:
            x: Input tensor.

        Returns:
            Output logits.
        )r  r   r   s     rO   r_   zResNetV2.forward  s)     !!!$a rP   )Tzresnet/F)N)NFFr  F)r   FT)&rd   re   rf   rg   rG   r|   r   r   r   r   rh   r
   r   r   ri   r   rA   rj   r   ignorer   r   r   r   r   r   Moduler   r	   r  rk   r   r  r  r  r   r_   rl   rm   s   @rO   r&   r&   	  sr    )?#$!# !""&"$''#*<B#G#,!$&#(-r9Ir9 CHor9 	r9
 r9 r9 r9 r9 r9 r9 r9 r9 r9  r9  r9  !!r9" !#r9$ %r9& "'r9( !)r9h YYQ4 Q4 Q Q YY5s 5C 5PT 5 5 YY	D 	T#s(^ 	 	 YY*T *T * *
 YY		  2C 2hsm 2W[ 2 8<$$',7 ||7  eCcN347  	7 
 7  7  !%7  
tELL!5tELL7I)I#JJ	K7 v ./$#	3S	>*  	 %,, 5<< 
Sell 
S 
S 
S %,, rP   r   namerW   c                 r   t        | t        j                        sd|v rpt        | t        j                        rVt        j                  j                  | j                  dd       t        j                  j                  | j                         yt        | t        j                        rct        j                  j                  | j                  dd       | j                  *t        j                  j                  | j                         yyt        | t        j                  t        j                  t        j                  f      rSt        j                  j                  | j                         t        j                  j                  | j                         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.
    head.fcra   g{Gz?)meanstdfan_outrelu)modenonlinearityNrW   )
isinstancerG   LinearConv2drS   normal_rU   rT   biaskaiming_normal_BatchNorm2d	LayerNorm	GroupNormones_hasattrrW   )r   r&  rW   s      rO   r   r     s
    &"))$d):z&RTR[R[?\
CT:
v{{#	FBII	&
IFS;;"GGNN6;;' #	FR^^R\\2<<H	I
fmm$
v{{#	GF,<= >rP   modelr   r   c                    dd l }d }|j                  |      }t        | j                  j                  j
                  j                  d    ||| d               }| j                  j                  j
                  j                  |       | j                  j
                  j                   ||| d                | j                  j                  j                   ||| d                t        t        | j                  dd       t        j                        r| j                  j                  j
                  j                  d   || d   j                  d	   k(  rv| j                  j                  j
                  j                   ||| d                | j                  j                  j                  j                   ||| d
                t!        | j"                  j%                               D ]]  \  }\  }}	t!        |	j&                  j%                               D ]-  \  }
\  }}d}| d|dz    d|
dz   dd}|j(                  j
                  j                   ||| d| d                |j*                  j
                  j                   ||| d| d                |j,                  j
                  j                   ||| d| d                |j.                  j
                  j                   ||| d                |j0                  j
                  j                   ||| d                |j2                  j
                  j                   ||| d                |j.                  j                  j                   ||| d                |j0                  j                  j                   ||| d                |j2                  j                  j                   ||| d                |j4                  || d| d   }|j4                  j                  j
                  j                   ||             0 ` y )Nr   c                 p    | j                   dk(  r| j                  g d      } t        j                  |       S )zPossibly convert HWIO to OIHW.r   )r>   r   r   r   )ndim	transposerj   
from_numpy)conv_weightss    rO   t2pz_load_weights.<locals>.t2p7  s1    !'11,?L--rP   r   z%root_block/standardized_conv2d/kernelzgroup_norm/gammazgroup_norm/betar   zhead/conv2d/kernelzhead/conv2d/biasstandardized_conv2dblockz/unit02d/za/z/kernelzb/zc/za/group_norm/gammazb/group_norm/gammazc/group_norm/gammaza/group_norm/betazb/group_norm/betazc/group_norm/betaza/proj/)numpyloadr"   r   r   rU   r  copy_r   r3  r/  r   r   rG   r1  r   r   r   named_childrenr   rD   rF   rr   rC   rE   rq   rB   )r:  r   r   nprA  weightsstem_conv_wisnamer   jbnamerD  cnameblock_prefixws                   rO   r   r   3  s   . ggo&G"

$$Q'Wx?d5e-f)giK	JJOO  -	JJC6(2B(C DEF	JJOO#g&@ABC'%**dD1299=JJMM  &&q)Wx?Q5R-S-Y-YZ\-]]

""3w&9K/L'M#NO

  Wx7G-H%I!JK&u||'B'B'DE ;>E5!*5<<+F+F+H!I 	;A~u)E$XU1q5'q1uSkCLKK$$SL>E7'1R)S%TUKK$$SL>E7'1R)S%TUKK$$SL>E7'1R)S%TUKK$$SL>AS1T)U%VWKK$$SL>AS1T)U%VWKK$$SL>AS1T)U%VWKK""3w,?P/Q'R#STKK""3w,?P/Q'R#STKK""3w,?P/Q'R#ST+|nGE7'BC  %%,,223q6:	;;rP   variant
pretrainedkwargsc                 B    t        d      }t        t        | |fd|i|S )zCreate a ResNetV2 model.

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

    Returns:
        ResNetV2 model instance.
    T)flatten_sequentialfeature_cfg)r   r   r&   )rU  rV  rW  rZ  s       rO   _create_resnetv2r[  Y  s4     $/K':  rP   c           	      @    t        | f|dt        t        d      d|S )zCreate a ResNetV2 model with BiT weights.

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

    Returns:
        ResNetV2 model instance.
    r   g:0yE>)eps)rV  r   r1   )r[  r   r   )rU  rV  rW  s      rO   _create_resnetv2_bitr^  l  s3     9$/	
  rP   urlc                 2    | dddddt         t        dddd	|S )
Nr   )r>      ra  )r   r   g      ?bilinearr   r(  z
apache-2.0)r_  r   
input_size	pool_sizecrop_pctinterpolationr)  r*  
first_conv
classifierlicenser   )r_  rW  s     rO   _cfgrj    s3    =vJ'0F!  rP   z%resnetv2_50x1_bit.goog_distilled_in1kztimm/bicubic)	hf_hub_idrf  custom_loadz-resnetv2_152x2_bit.goog_teacher_in21k_ft_in1kz1resnetv2_152x2_bit.goog_teacher_in21k_ft_in1k_384)r>     rn  )   ro  r`   )rl  rc  rd  re  rf  rm  z$resnetv2_50x1_bit.goog_in21k_ft_in1k)r>     rp  )   rq  )rl  rc  rd  re  rm  z$resnetv2_50x3_bit.goog_in21k_ft_in1kz%resnetv2_101x1_bit.goog_in21k_ft_in1kz%resnetv2_101x3_bit.goog_in21k_ft_in1kz%resnetv2_152x2_bit.goog_in21k_ft_in1kz%resnetv2_152x4_bit.goog_in21k_ft_in1k)r>     rr  )   rs  zresnetv2_50x1_bit.goog_in21kiSU  )rl  r   rm  zresnetv2_50x3_bit.goog_in21kzresnetv2_101x1_bit.goog_in21kzresnetv2_101x3_bit.goog_in21kzresnetv2_152x2_bit.goog_in21kzresnetv2_152x4_bit.goog_in21kzresnetv2_18.ra4_e3600_r224_in1kg?)r>      rt  )rl  rf  re  test_input_sizetest_crop_pctz resnetv2_18d.ra4_e3600_r224_in1kz
stem.conv1)rl  rf  re  ru  rv  rg  )rl  re  rc  rd  ru  rf  rg  gffffff?)rf  rg  )rf  )rl  rf  rg  re  ru  rv  )zresnetv2_34.ra4_e3600_r224_in1kz resnetv2_34d.ra4_e3600_r224_in1kz resnetv2_34d.ra4_e3600_r384_in1kzresnetv2_50.a1h_in1kzresnetv2_50d.untrainedzresnetv2_50t.untrainedzresnetv2_101.a1h_in1kzresnetv2_101d.untrainedzresnetv2_152.untrainedzresnetv2_152d.untrainedzresnetv2_50d_gn.ah_in1kzresnetv2_50d_evos.ah_in1kzresnetv2_50d_frn.untrainedc                 &    t        	 d| g ddd|S )zResNetV2-50x1-BiT model.r>   r      r>   r   rV  r   r   )resnetv2_50x1_bitr^  rV  rW  s     rO   r{  r{    -      c(2<VWc[ac crP   c                 &    t        	 d| g ddd|S )zResNetV2-50x3-BiT model.rx  r>   rz  )resnetv2_50x3_bitr|  r}  s     rO   r  r    r~  rP   c                 &    t        	 d| g ddd|S )zResNetV2-101x1-BiT model.r>   r      r>   r   rz  )resnetv2_101x1_bitr|  r}  s     rO   r  r    -      e)3MXYe]ce erP   c                 &    t        	 d| g ddd|S )zResNetV2-101x3-BiT model.r  r>   rz  )resnetv2_101x3_bitr|  r}  s     rO   r  r    r  rP   c                 &    t        	 d| g ddd|S )zResNetV2-152x2-BiT model.r>   r   $   r>   r   rz  )resnetv2_152x2_bitr|  r}  s     rO   r  r    r  rP   c                 &    t        	 d| g ddd|S )zResNetV2-152x4-BiT model.r  r   rz  )resnetv2_152x4_bitr|  r}  s     rO   r  r    r  rP   c           	      f    t        g ddddt        t              }t        dd| it        |fi |S )zResNetV2-18 model.r   r   r   r   r      r   r   Tr`   r   r   r   r+   r1   r2   rV  )resnetv2_18r   r   r   r[  rV  rW  
model_argss      rO   r  r    s>     &9TW ^J _j_DD^W]D^__rP   c           
      j    t        g ddddt        t        dd      }t        dd| it        |fi |S )	z'ResNetV2-18d model (deep stem variant).r  r  Tr`   r   r   r   r   r+   r1   r2   r   r   rV  )resnetv2_18dr  r  s      rO   r  r  $  sC     &9TW ^vX\J `z`T*E_X^E_``rP   c           	      b    t        ddddt        t              }t        dd| it        |fi |S )zResNetV2-34 model.rx  r  Tr`   r  rV  )resnetv2_34r  r  s      rO   r  r  .  s>     &9TW ^J _j_DD^W]D^__rP   c           
      f    t        ddddt        t        dd      }t        dd| it        |fi |S )	z'ResNetV2-34d model (deep stem variant).rx  r  Tr`   r   r  rV  )resnetv2_34dr  r  s      rO   r  r  8  sC     &9TW ^vX\J `z`T*E_X^E_``rP   c           	      `    t        g dt        t              }t        dd| it        |fi |S )zResNetV2-50 model.rx  r   r1   r2   rV  )resnetv2_50r  r  s      rO   r  r  B  s1     \mP^_J_j_DD^W]D^__rP   c           	      d    t        g dt        t        dd      }t        dd| it        |fi |S )z'ResNetV2-50d model (deep stem variant).rx  r   Tr   r1   r2   r   r   rV  )resnetv2_50dr  r  s      rO   r  r  I  s:     .4)J `z`T*E_X^E_``rP   c           	      d    t        g dt        t        dd      }t        dd| it        |fi |S )z)ResNetV2-50t model (tiered stem variant).rx  r   Tr  rV  )resnetv2_50tr  r  s      rO   r  r  R  s:     .T+J `z`T*E_X^E_``rP   c           	      `    t        g dt        t              }t        dd| it        |fi |S )zResNetV2-101 model.r  r  rV  )resnetv2_101r  r  s      rO   r  r  [  1     ]}Q_`J`z`T*E_X^E_``rP   c           	      d    t        g dt        t        dd      }t        dd| it        |fi |S )z(ResNetV2-101d model (deep stem variant).r  r   Tr  rV  )resnetv2_101dr  r  s      rO   r  r  b  :     >4)J a
ad:F`Y_F`aarP   c           	      `    t        g dt        t              }t        dd| it        |fi |S )zResNetV2-152 model.r  r  rV  )resnetv2_152r  r  s      rO   r  r  k  r  rP   c           	      d    t        g dt        t        dd      }t        dd| it        |fi |S )z(ResNetV2-152d model (deep stem variant).r  r   Tr  rV  )resnetv2_152dr  r  s      rO   r  r  r  r  rP   c           	      d    t        g dt        t        dd      }t        dd| it        |fi |S )z,ResNetV2-50d model with Group Normalization.rx  r   Tr  rV  )resnetv2_50d_gn)r   r   r   r[  r  s      rO   r  r  }  s:     ,4)J c*cZHb[aHbccrP   c           	      d    t        g dt        t        dd      }t        dd| it        |fi |S )z ResNetV2-50d model with EvoNorm.rx  r   Tr  rV  )resnetv2_50d_evos)r   r   r   r[  r  s      rO   r  r    s:     +4)J eJe$zJd]cJdeerP   c           	      d    t        g dt        t        dd      }t        dd| it        |fi |S )z6ResNetV2-50d model with Filter Response Normalization.rx  r   Tr  rV  )resnetv2_50d_frn)r   r   r   r[  r  s      rO   r  r    s;     BY4)J d:djIc\bIcddrP   )resnetv2_50x1_bitmresnetv2_50x3_bitmresnetv2_101x1_bitmresnetv2_101x3_bitmresnetv2_152x2_bitmresnetv2_152x4_bitmresnetv2_50x1_bitm_in21kresnetv2_50x3_bitm_in21kresnetv2_101x1_bitm_in21kresnetv2_101x3_bitm_in21kresnetv2_152x2_bitm_in21kresnetv2_152x4_bitm_in21kresnetv2_50x1_bit_distilledresnetv2_152x2_bit_teacherresnetv2_152x2_bit_teacher_384)r   Tr"  r#  )r   )Yrg   collectionsr   	functoolsr   typingr   r   r   r   r	   r
   r   rj   torch.nnrG   	timm.datar   r   timm.layersr   r   r   r   r   r   r   r   r   r   r   r   r   r   _builderr   	_featuresr   _manipulater    r!   r"   	_registryr#   r$   r%   __all__r%  r(   ro   rx   r   r   r   r   r   r   rh   r   r   r&   r   no_gradr   r[  r^  rj  default_cfgsr{  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  rd    rP   rO   <module>r     s		  > $  D D D   ES S S S * + F F Y Y,Y")) Yxbryy bJP Pf'RYY 'B$2BII $2N<")) <~	=C 	=D 	= (&|C--- - 	-
 - - ]]-`Qryy Qh "))  3  T  UY  , "; ";S ";# "; ";Jc t s x &# 4 3 S[ (	c 	# 	$sCx. 	 % a&+TT.3a&
 4TT63a& 8 HsR[im:oa& +D HsPT-Va&  +D HsPT-V!a&& ,T HsPT.V'a&, ,T HsPT.V-a&2 ,T HsPT.V3a&8 ,T HsPT.V9a&B #Dt%-Ca&H #Dt%-Ia&N $Tt&-Oa&T $Tt&-Ua&Z $Tt&-[a&` $Tt&-aa&h &t#}\_(aia&n '#}\_)!oa&v (,#}\_(a )-#}\_)! )-(TaL): !$]`b #L:"L:!$]`b  $L :"!#L :  $L}C I "&L}C"I #'L#:a& aH c$ c# c( c c c$ c# c( c c e4 e3 e8 e e e4 e3 e8 e e e4 e3 e8 e e e4 e3 e8 e e `D `C `H ` ` aT aS aX a a `D `C `H ` ` aT aS aX a a `D `C `H ` ` aT aS aX a a aT aS aX a a aT aS aX a a bd bc bh b b aT aS aX a a bd bc bh b b d d d d d f$ f# f( f f e e e e e H@@BBBB > >!@!@!@!@#J"Q&Y' rP   