
    ^jd              
       r   d Z ddlZddlZddlZddlmZ ddlmZm	Z	m
Z
mZmZ ddlZddlmc m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! dd	l"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gZ/ ej`                  e1      Z2 G d dejf                        Z4 G d dejf                        Z5 G d dejf                        Z6de7fdZ8e'de7fd       Z9 G d dejf                        Z: G d dejf                        Z;d1dejf                  de<de=fdZ>d Z?d  Z@d2d!ZAd3d"ZB e- eB        eB        eB        eBd#$       eBd#$       eBd#$      d%      ZCe,d2d&e;fd'       ZDe,d2d&e;fd(       ZEe,d2d&e;fd)       ZFe,d2d&e;fd*       ZGe,d2d&e;fd+       ZHe,d2d&e;fd,       ZI e.e1d-d.d/d0       y)4a   Nested Transformer (NesT) in PyTorch

A PyTorch implement of Aggregating Nested Transformers as described in:

'Aggregating Nested Transformers'
    - https://arxiv.org/abs/2105.12723

The official Jax code is released and available at https://github.com/google-research/nested-transformer. The weights
have been converted with convert/convert_nest_flax.py

Acknowledgments:
* The paper authors for sharing their research, code, and model weights
* Ross Wightman's existing code off which I based this

Copyright 2021 Alexander Soare
    N)partial)ListOptionalTupleTypeUnion)nnIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)
PatchEmbedMlpDropPathcalculate_drop_path_ratescreate_classifiertrunc_normal__assertcreate_conv2dcreate_pool2d	to_ntupleuse_fused_attn	LayerNorm   )build_model_with_cfg)feature_take_indices)register_notrace_function)checkpoint_seqnamed_apply)register_modelgenerate_default_cfgsregister_model_deprecationsNestc                        e Zd ZU dZej
                  j                  e   ed<   	 	 	 	 	 	 d
de	de	dede
de
f
 fdZd	 Z xZS )	Attentionz
    This is much like `.vision_transformer.Attention` but uses *localised* self attention by accepting an input with
     an extra "image block" dim
    
fused_attndim	num_headsqkv_bias	attn_drop	proj_dropc                 X   ||d}t         
|           || _        ||z  }	|	dz  | _        t	               | _        t        j                  |d|z  fd|i|| _        t        j                  |      | _
        t        j                  ||fi || _        t        j                  |      | _        y )Ndevicedtypeg         bias)super__init__r'   scaler   r%   r	   LinearqkvDropoutr)   projr*   )selfr&   r'   r(   r)   r*   r-   r.   ddhead_dim	__class__s             [/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/nest.pyr2   zAttention.__init__=   s     /")#%
(*99S!C%=h="=I.IIc3-"-	I.    c           	         |j                   \  }}}}| j                  |      j                  |||d| j                  || j                  z        j	                  dddddd      }|j                  d      \  }}}	| j                  r<t        j                  |||	| j                  r| j                  j                  nd      }nL|| j                  z  }||j                  d	d
      z  }
|
j                  d
      }
| j                  |
      }
|
|	z  }|j	                  ddddd      j                  ||||      }| j                  |      }| j!                  |      }|S )zm
        x is shape: B (batch_size), T (image blocks), N (seq length per image block), C (embed dim)
        r/   r      r                 )	dropout_p)r&   )shaper5   reshaper'   permuteunbindr%   Fscaled_dot_product_attentiontrainingr)   pr3   	transposesoftmaxr7   r*   )r8   xBTNCr5   qkvattns              r<   forwardzAttention.forwardS   s7    WW
1ahhqk!!!Q1dnna4>>>QRZZ[\^_abdeghjkl**Q-1a??..q!QVZVcVc$..BRBRiklADJJAq{{2r**D<<B<'D>>$'DqA IIaAq!$,,Q1a8IIaLNN1r=   )   FrB   rB   NN)__name__
__module____qualname____doc__torchjitFinalbool__annotations__intfloatr2   rY   __classcell__r;   s   @r<   r$   r$   6   sk     		%%
 "!!// / 	/
 / /,r=   r$   c                        e Zd ZdZdddddej
                  ej                  ddf	dededed	e	d
ededede
ej                     de
ej                     f fdZd Z xZS )TransformerLayerz
    This is much like `.vision_transformer.Block` but:
        - Called TransformerLayer here to allow for "block" as defined in the paper ("non-overlapping image blocks")
        - Uses modified Attention layer that handles the "block" dimension
          @FrB   Nr&   r'   	mlp_ratior(   r*   r)   	drop_path	act_layer
norm_layerc                    |
|d}t         |            |	|fi || _        t        |f||||d|| _        |dkD  rt        |      nt        j                         | _         |	|fi || _	        t        ||z        }t        d||||d|| _        |dkD  rt        |      | _        y t        j                         | _        y )Nr,   )r'   r(   r)   r*   rB   )in_featureshidden_featuresrm   drop )r1   r2   norm1r$   rX   r   r	   Identity
drop_path1norm2rd   r   mlp
drop_path2)r8   r&   r'   rk   r(   r*   r)   rl   rm   rn   r-   r.   r9   mlp_hidden_dimr;   s                 r<   r2   zTransformerLayer.__init__r   s     /*r*


 
	 2;R(9-R[[]*r*
S9_- 
*	

 
 2;R(9-R[[]r=   c                     | j                  |      }|| j                  | j                  |            z   }|| j                  | j	                  | j                  |                  z   }|S N)rt   rv   rX   ry   rx   rw   )r8   rP   ys      r<   rY   zTransformerLayer.forward   sS    JJqM		!--A 788r=   )r[   r\   r]   r^   r	   GELUr   rd   re   rb   r   Moduler2   rY   rf   rg   s   @r<   ri   ri   l   s      ""!!!)+*,,,#S#S #S 	#S
 #S #S #S #S BII#S RYY#SJr=   ri   c            	       X     e Zd Z	 	 	 ddededeej                     def fdZd Z	 xZ
S )ConvPoolin_channelsout_channelsrn   pad_typec                     ||d}t         |           t        ||fd|dd|| _         ||fi || _        t        ddd|      | _        y )Nr,   r/   T)kernel_sizepaddingr0   maxr@   )r   strider   )r1   r2   r   convnormr   pool)	r8   r   r   rn   r   r-   r.   r9   r;   s	           r<   r2   zConvPool.__init__   s\     /!+|nT\cgnkmn	|2r2	!%Qq(S	r=   c                 0   t        |j                  d   dz  dk(  d       t        |j                  d   dz  dk(  d       | j                  |      }| j                  |j	                  dddd            j	                  dddd      }| j                  |      }|S )z:
        x is expected to have shape (B, C, H, W)
        rD   r@   r   z1BlockAggregation requires even input spatial dimsrE   r/   r   )r   rF   r   r   rH   r   r8   rP   s     r<   rY   zConvPool.forward   s     	a1$&YZa1$&YZIIaLIIaii1a+,44Q1a@IIaLr=   ) NN)r[   r\   r]   rd   r   r	   r   strr2   rY   rf   rg   s   @r<   r   r      sK     TT T RYY	T
 T
r=   r   
block_sizec                     | j                   \  }}}}t        ||z  dk(  d       t        ||z  dk(  d       ||z  }||z  }| j                  ||||||      } | j                  dd      j                  |||z  d|      } | S )zimage to blocks
    Args:
        x (Tensor): with shape (B, H, W, C)
        block_size (int): edge length of a single square block in units of H, W
    r   z,`block_size` must divide input height evenlyz+`block_size` must divide input width evenlyr@   r/   rE   )rF   r   rG   rN   )rP   r   rQ   HWrT   grid_height
grid_widths           r<   blockifyr      s     ''JAq!QA
Na!OPA
Na!NOz/KjJ			![*j*aHA	Aq!!![:%=r1EAHr=   c                     | j                   \  }}}}t        t        j                  |            }||z  x}}| j	                  ||||||      } | j                  dd      j	                  ||||      } | S )zblocks to image
    Args:
        x (Tensor): with shape (B, T, N, C) where T is number of blocks and N is sequence size per block
        block_size (int): edge length of a single square block in units of desired H, W
    r@   r/   )rF   rd   mathsqrtrG   rN   )	rP   r   rQ   rR   _rT   	grid_sizeheightwidths	            r<   
deblockifyr      st     JAq!QDIIaL!I++FU			!Y	:z1EA	Aq!!!VUA6AHr=   c                        e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 ddededededededee   d	ed
edededeee      dee	e
j                        dee	e
j                        def fdZd Z xZS )	NestLevelz7 Single hierarchical level of a Nested Transformer
    
num_blocksr   
seq_lengthr'   depth	embed_dimprev_embed_dimrk   r(   r*   r)   rl   rn   rm   r   c                    ||d}t         |           || _        d| _        t	        j
                  t        j                  d|||fi |      | _        |t        ||f||d|| _
        nt	        j                         | _
        t        |      rt        |      |k(  sJ d       t	        j                  t        |      D cg c]  }t        d||||	|
||r||   nd ||d	|  c} | _        y c c}w )Nr,   Fr   )rn   r   zDMust provide as many drop path rates as there are transformer layers)	r&   r'   rk   r(   r*   r)   rl   rn   rm   rs   )r1   r2   r   grad_checkpointingr	   	Parameterr_   zeros	pos_embedr   r   ru   len
Sequentialrangeri   transformer_encoder)r8   r   r   r   r'   r   r   r   rk   r(   r*   r)   rl   rn   rm   r   r-   r.   r9   ir;   s                       r<   r2   zNestLevel.__init__   s	   ( /$"'ekk!ZY&]Z\&]^% kz\dkhjkDIDI y>y>U*r,rr*#%== 5\3#   ##!##*3)A,%# 3# $$  3#s   #C.c                    | j                  |      }|j                  dddd      }t        || j                        }|| j                  z   }| j
                  r5t        j                  j                         st        | j                  |      }n| j                  |      }t        || j                        }|j                  dddd      S )z+
        expects x as (B, C, H, W)
        r   r@   r/   r   )r   rH   r   r   r   r   r_   r`   is_scriptingr   r   r   r   s     r<   rY   zNestLevel.forward  s     IIaLIIaAq!Q(""599+A+A+Ct77;A((+Aq$//*yyAq!$$r=   )Nrj   TrB   rB   NNNr   NN)r[   r\   r]   r^   rd   r   re   rb   r   r   r	   r   r   r2   rY   rf   rg   s   @r<   r   r      s     -1!!!!/34837%0$0$ 0$ 	0$
 0$ 0$ 0$ %SM0$ 0$ 0$ 0$ 0$  U,0$ !bii10$  RYY00$  !0$d%r=   r   c            '       @    e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d,dededededeedf   deedf   d	eedf   d
edededededededee	e
j                        dee	e
j                        dededef& fdZej                  j                   d-d       Zej                  j                   d        Zej                  j                   d.d       Zej                  j                   d/d       Zej                  j                   de
j                  fd       Zd0d
edefdZ	 	 	 	 	 d1dej.                  deeeee   f      d ed!ed"ed#edeeej.                     eej.                  eej.                     f   f   fd$Z	 	 	 d2deeee   f   d%ed&efd'Zd( Zd.d)efd*Zd+ Z xZS )3r"   z Nested Transformer (NesT)

    A PyTorch impl of : `Aggregating Nested Transformers`
        - https://arxiv.org/abs/2105.12723
    img_sizein_chans
patch_size
num_levels
embed_dims.r'   depthsnum_classesrk   r(   	drop_rateproj_drop_rateattn_drop_ratedrop_path_ratern   rm   r   weight_initglobal_poolc                    t          |           ||d}dD ]M  }t               |   }t        |t        j
                  j                        s5t        |      |k(  rDJ d| d         t        |      |      } t        |      |      } t        |      |      }|| _	        || _
        |d   x| _        | _        g | _        |xs t        }|xs t        j                   }|| _        || _        t        |t        j
                  j                        r|d   |d   k(  sJ d       |d   }||z  dk(  sJ d	       || _        d
t)        j*                  |dt(        j,                        z  j/                  d      j1                         | _        ||z  t5        j6                  | j2                  d         z  dk(  sJ d       t9        ||z  t5        j6                  | j2                  d         z        | _        t=        d||||d   dd|| _        | j>                  j@                  | _         | j@                  | j2                  d   z  | _!        g }tE        ||d      }d}d
}tG        t        | j2                              D ]  }||   }|jI                  tK        | j2                  |   | j:                  | jB                  ||   ||   ||f|	|
||||   |||d|       | xj                  tM        ||d|       gz  c_        |}|dz  } t        jN                  | | _(         ||d   fi || _)        tU        | j                  | j                  fd|i|\  }}|| _+        t        jX                  |      | _-        || _.        | j_                  |       y)a  
        Args:
            img_size (int, tuple): input image size
            in_chans (int): number of input channels
            patch_size (int): patch size
            num_levels (int): number of block hierarchies (T_d in the paper)
            embed_dims (int, tuple): embedding dimensions of each level
            num_heads (int, tuple): number of attention heads for each level
            depths (int, tuple): number of transformer layers for each level
            num_classes (int): number of classes for classification head
            mlp_ratio (int): ratio of mlp hidden dim to embedding dim for MLP of transformer layers
            qkv_bias (bool): enable bias for qkv if True
            drop_rate (float): dropout rate for MLP of transformer layers, MSA final projection layer, and classifier
            attn_drop_rate (float): attention dropout rate
            drop_path_rate (float): stochastic depth rate
            norm_layer: (nn.Module): normalization layer for transformer layers
            act_layer: (nn.Module): activation layer in MLP of transformer layers
            pad_type: str: Type of padding to use '' for PyTorch symmetric, 'same' for TF SAME
            weight_init: (str): weight init scheme
            global_pool: (str): type of pooling operation to apply to final feature map

        Notes:
            - Default values follow NesT-B from the original Jax code.
            - `embed_dims`, `num_heads`, `depths` should be ints or tuples with length `num_levels`.
            - For those following the paper, Table A1 may have errors!
                - https://github.com/google-research/nested-transformer/issues/2
        r,   r   r'   r   zRequire `len(z) == num_levels`rE   r   r   z Model only handles square inputsz*`patch_size` must divide `img_size` evenlyr?   cpuzUFirst level blocks don't fit evenly. Check `img_size`, `patch_size`, and `num_levels`F)r   r   r   r   flattenT)	stagewiseN)rk   r(   r*   r)   rl   rn   rm   r   zlevels.)num_chs	reductionmoduler@   	pool_typers   )0r1   r2   locals
isinstancecollectionsabcSequencer   r   r   r   num_featureshead_hidden_sizefeature_infor   r	   r~   r   r   r   r_   arangelongfliptolistr   r   r   rd   r   r   patch_embednum_patchesr   r   r   appendr   dictr   levelsr   r   r   r6   	head_dropheadinit_weights)!r8   r   r   r   r   r   r'   r   r   rk   r(   r   r   r   r   rn   rm   r   r   r   r-   r.   r9   
param_nameparam_valuer   dp_ratesprev_dimcurr_strider   r&   r   r;   s!                                   r<   r2   zNest.__init__'  s   f 	/? 	dJ (:.K+{'?'?@;':5czlRb7cc5	d
 +Yz*:6
)Ij))4	&:&v.& 4>rNBD1,9
(	"$h 8 89A;(1+-Q/QQ-{H*$)W+WW)$ ZUZZ XX^^_`ahhjJ&$))DOOA4F*GG1L 	ed	eL x:5$))DOOTUDV:WWX & 
! m
 
  ++77**dooa.@@ ,^VtTs4??+, 	AQ-CMM)"!q	 $!(("1+%#!  ! $ $skT[\][^R_"`!aaH1K-	. mmV, z"~44	 .d.?.?AQAQo]holnoT&I.	+&r=   c                     |dv sJ d|v r t        j                  | j                         nd}| j                  D ]  }t	        |j
                  ddd        t        t        t        |      |        y )	N)nlhbr   r   rB   {Gz?rD   r@   stdab)	head_bias)	r   logr   r   r   r   r   r   _init_nest_weights)r8   moder   levels       r<   r   zNest.init_weights  sf    |###39T>TXXd..//r	[[ 	?E%//sbA>	?G.)DdKr=   c                 l    t        t        | j                              D ch c]  }d| d
 c}S c c}w )Nzlevel.z
.pos_embed)r   r   r   )r8   r   s     r<   no_weight_decayzNest.no_weight_decay  s-    05c$++6F0GH1&:&HHHs   1c                 2    t        d|rdndd fddg      }|S )Nz^patch_embedz^levels\.(\d+)z*^levels\.(\d+)\.transformer_encoder\.(\d+))z"^levels\.(\d+)\.(?:pool|pos_embed))r   )z^norm)i )stemblocks)r   )r8   coarsematchers      r<   group_matcherzNest.group_matcher  s0     &,"2_aef=$
 r=   c                 4    | j                   D ]	  }||_         y r|   )r   r   )r8   enablels      r<   set_grad_checkpointingzNest.set_grad_checkpointing  s     	*A#)A 	*r=   returnc                     | j                   S r|   )r   )r8   s    r<   get_classifierzNest.get_classifier  s    yyr=   c                 p    || _         t        | j                  | j                   |      \  | _        | _        y )N)r   )r   r   r   r   r   )r8   r   r   s      r<   reset_classifierzNest.reset_classifier  s2    &&7t//;'H#$)r=   rP   indicesr   
stop_early
output_fmtintermediates_onlyc           	         |dv sJ d       g }t        t        | j                        |      \  }}	| j                  |      }t        | j                        dz
  }
t
        j                  j                         s|s| j                  }n| j                  d|	dz    }t        |      D ]q  \  }} ||      }||v s|rL||
k(  rG| j                  |j                  dddd            j                  dddd      }|j                  |       a|j                  |       s |r|S |
k(  r5| j                  |j                  dddd            j                  dddd      }||fS )a   Forward features that returns intermediates.

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

        )NCHWzOutput shape must be NCHW.r   Nr   r@   r/   )r   r   r   r   r   r_   r`   r   	enumerater   rH   r   )r8   rP   r   r   r  r  r  intermediatestake_indices	max_indexlast_idxstagesfeat_idxstagex_inters                  r<   forward_intermediateszNest.forward_intermediates  sU   * Y&D(DD&"6s4;;7G"Qi Qt'!+99!!#:[[F[[)a-0F(0 	,OHeaA<'H0"ii		!Q1(=>FFq!QPQRG!((1!((+	,   x		!))Aq!Q/088Aq!DA-r=   
prune_norm
prune_headc                     t        t        | j                        |      \  }}| j                  d|dz    | _        |rt        j                         | _        |r| j                  dd       |S )z@ Prune layers not required for specified intermediates.
        Nr   r   r   )r   r   r   r	   ru   r   r   )r8   r   r  r  r  r	  s         r<   prune_intermediate_layerszNest.prune_intermediate_layers
  s]     #7s4;;7G"Qikk.9q=1DI!!!R(r=   c                     | j                  |      }| j                  |      }| j                  |j                  dddd            j                  dddd      }|S )Nr   r@   r/   r   )r   r   r   rH   r   s     r<   forward_featureszNest.forward_features  sR    QKKNIIaii1a+,44Q1a@r=   
pre_logitsc                 p    | j                  |      }| j                  |      }|r|S | j                  |      S r|   )r   r   r   )r8   rP   r  s      r<   forward_headzNest.forward_head!  s5    QNN1q0DIIaL0r=   c                 J    | j                  |      }| j                  |      }|S r|   )r  r  r   s     r<   rY   zNest.forward&  s'    !!!$a r=   )   r/   r?   r/         i   r?   rZ      r@   r@        rj   TrB   rB   rB   g      ?NNr   r   avgNNr   F)T)r#  )NFFr  F)r   FT) r[   r\   r]   r^   rd   r   re   rb   r   r   r	   r   r   r2   r_   r`   ignorer   r   r   r   r   r   Tensorr   r   r  r  r  r  rY   rf   rg   s   @r<   r"   r"      s     *9)3&0#!!!$&$&$'4837!$-H'H' H' 	H'
 H' c3hH' S#XH' #s(OH' H' H' H' H' "H' "H' "H'  !bii1!H'"  RYY0#H'$ %H'& 'H'( )H'T YYL L YYI I YY	 	 YY* * YY		  HC Hc H 8<$$',1 ||1  eCcN341  	1 
 1  1  !%1  
tELL!5tELL7I)I#JJ	K1 j ./$#	3S	>*  	 1$ 1
r=   r   namer   c                 V   t        | t        j                        r|j                  d      rDt	        | j
                  ddd       t        j                  j                  | j                  |       yt	        | j
                  ddd       | j                  *t        j                  j                  | j                         yyt        | t        j                        rPt	        | j
                  ddd       | j                  *t        j                  j                  | j                         yyy)zn NesT weight initialization
    Can replicate Jax implementation. Otherwise follows vision_transformer.py
    r   r   rD   r@   r   N)r   r	   r4   
startswithr   weightinit	constant_r0   zeros_Conv2d)r   r(  r   s      r<   r   r   ,  s     &"))$??6"&--SB!<GGfkk95&--SB!<{{&v{{+ '	FBII	&fmma8;;"GGNN6;;' # 
'r=   c                    t         j                  d| j                  |j                         | j                  d   }|j                  dd \  }}t        t	        j
                  ||z              }t        | t        t	        j
                  |                  j                  dddd      } t        j                  | ||gdd      } t        | j                  dddd      t        t	        j
                  |                  } | S )	z
    Rescale the grid of position embeddings when loading from state_dict
    Expected shape of position embeddings is (1, T, N, C), and considers only square images
    z$Resized position embedding: %s to %sr@   r   r/   r   bicubicF)sizer   align_corners)_loggerinforF   rd   r   r   r   rH   rJ   interpolater   )posemb
posemb_newseq_length_oldnum_blocks_newseq_length_newsize_news         r<   resize_pos_embedr=  >  s    
 LL7zGWGWX\\!_N%/%5%5a%:"NN499^N:;<HDIIn$= >?GG1aQRSF]]68(<9\abFfnnQ1a0#dii6O2PQFMr=   c                    | j                         D cg c]  }|j                  d      s| }}|D ]E  }| |   j                  t        ||      j                  k7  s*t	        | |   t        ||            | |<   G | S c c}w )z4 resize positional embeddings of pretrained weights 
pos_embed_)keysr*  rF   getattrr=  )
state_dictmodelrV   pos_embed_keyss       r<   checkpoint_filter_fnrE  O  s    !+!2QAall<6PaQNQ Oa='%"3"9"99,Z]GE1<MNJqMO 	 Rs
   A<A<c                 N    t        t        | |ft        dd      t        d|}|S )N)r   r   r@   T)out_indicesflatten_sequential)feature_cfgpretrained_filter_fn)r   r"   r   rE  )variant
pretrainedkwargsrC  s       r<   _create_nestrN  X  s:      Y4H1 E Lr=   c                 8    | ddddgdddt         t        ddd	d
|S )Nr"  )r/   r  r     g      ?r1  Tzpatch_embed.projr   z
apache-2.0)urlr   
input_size	pool_sizecrop_pctinterpolationfixed_input_sizemeanr   
first_conv
classifierlicenser
   )rQ  rM  s     r<   _cfgr[  e  s9    =Bx9$%.B(  r=   ztimm/)	hf_hub_id)znest_base.untrainedznest_small.untrainedznest_tiny.untrainedznest_base_jx.goog_in1kznest_small_jx.goog_in1kznest_tiny_jx.goog_in1kr   c                 >    t        ddddd|}t        dd| i|}|S ) Nest-B @ 224x224
    r  r  r   r   rL  rs   )	nest_baser   rN  rL  rM  model_kwargsrC  s       r<   r_  r_  |  s<      W"jWOUWLLL|LELr=   c                 >    t        ddddd|}t        dd| i|}|S ) Nest-S @ 224x224
    `      i  r/         r   r   rL  rs   )
nest_smallr`  ra  s       r<   rk  rk    s3     e>ZPZe^deLM*MMELr=   c                 >    t        ddddd|}t        dd| i|}|S ) Nest-T @ 224x224
    re  rh  r@   r@   rZ   r   rL  rs   )	nest_tinyr`  ra  s       r<   ro  ro    s3     d>ZPYd]cdLLL|LELr=   c                 b    |j                  dd       t        ddddd|}t        d	d| i|}|S )
r^  r   samer  r  r   r   rL  rs   )nest_base_jx
setdefaultr   rN  ra  s       r<   rr  rr    sL     j&) W"jWOUWLOJO,OELr=   c                 b    |j                  dd       t        ddddd|}t        d	d| i|}|S )
rd  r   rq  re  rh  r   r   rL  rs   )nest_small_jxrs  ra  s       r<   rv  rv    sC     j&)e>ZPZe^deLPZP<PELr=   c                 b    |j                  dd       t        ddddd|}t        d	d| i|}|S )
rm  r   rq  re  rh  rn  r   rL  rs   )nest_tiny_jxrs  ra  s       r<   rx  rx    sC     j&)d>ZPYd]cdLOJO,OELr=   rr  rv  rx  )jx_nest_basejx_nest_smalljx_nest_tiny)r   rB   r%  r$  )Jr^   collections.abcr   loggingr   	functoolsr   typingr   r   r   r   r   r_   torch.nn.functionalr	   
functionalrJ   	timm.datar   r   timm.layersr   r   r   r   r   r   r   r   r   r   r   r   _builderr   	_featuresr   _features_fxr   _manipulater   r   	_registryr   r    r!   __all__	getLoggerr[   r4  r   r$   ri   r   rd   r   r   r   r"   r   re   r   r=  rE  rN  r[  default_cfgsr_  rk  ro  rr  rv  rx  rs   r=   r<   <module>r     s  "     5 5     A    + + 3 4 Y Y(
'

H
%3		 3l/ryy /dryy :C   c  C%		 C%LI299 IX(ryy ( (U ($"
	 %6 F6"W5#g6"W5&  T   d   T            H"$"' r=   