
    ^jR                        d Z ddlZddlmZ ddl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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 dd	lmZ dd
lmZmZ ddl m!Z!m"Z"m#Z# dgZ$ G d dejJ                        Z& G d dejN                        Z( G d dejR                        Z* G d dejJ                        Z+de,dee-ej\                  f   fdZ/de-de0dee0df   de1de+f
dZ2d-de-dee-ef   fdZ3 e" e3dd d!"       e3dd d!"       e3        e3d#       e3d#       e3d#       e3d#      d$      Z4e!d.de+fd%       Z5e!d.de+fd&       Z6e!d.de+fd'       Z7e!d.de+fd(       Z8e!d.de+fd)       Z9e!d.de+fd*       Z: e#e;d+d,i       y)/zPytorch Densenet implementation w/ tweaks
This file is a copy of https://github.com/pytorch/vision 'densenet.py' (BSD-3-Clause) with
fixed kwargs passthrough and addition of dynamic global avg/max pool.
    N)OrderedDict)AnyDictOptionalTupleTypeUnion)ListIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)BatchNormAct2dget_norm_act_layer
BlurPool2dcreate_classifier   )build_model_with_cfg)MATCH_PREV_GROUP
checkpoint)register_modelgenerate_default_cfgsregister_model_deprecationsDenseNetc                   h    e Zd ZdZeddddfdedededeej                     d	e	d
e
ddf fdZdeej                     dej                  fdZdeej                     de
fdZdeej                     dej                  fdZdeej                  eej                     f   dej                  fdZ xZS )
DenseLayerzbDense layer for DenseNet.

    Implements the bottleneck layer with 1x1 and 3x3 convolutions.
            FNnum_input_featuresgrowth_ratebn_size
norm_layer	drop_rategrad_checkpointingreturnc	                    ||d}	t         
|           | j                  d ||fi |	      f | j                  dt        j                  |||z  fdddd|	      f | j                  d |||z  fi |	      f | j                  dt        j                  ||z  |fd	dddd
|	      f t        |      | _        || _        y)ad  Initialize DenseLayer.

        Args:
            num_input_features: Number of input features.
            growth_rate: Growth rate (k) of the layer.
            bn_size: Bottleneck size multiplier.
            norm_layer: Normalization layer class.
            drop_rate: Dropout rate.
            grad_checkpointing: Use gradient checkpointing.
        devicedtypenorm1conv1r   Fkernel_sizestridebiasnorm2conv2   r+   r,   paddingr-   N)super__init__
add_modulennConv2dfloatr!   r"   )selfr   r   r   r    r!   r"   r&   r'   dd	__class__s             _/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/densenet.pyr4   zDenseLayer.__init__   s    * /,>!E"!EFG+ 5"bCDQUZ"b^`"b 	c 	dGk,A!HR!HIJk!;"f<=aQRY^"fbd"f 	g 	hy)"4    xsc                 r    t        j                  |d      }| j                  | j                  |            }|S )z.Bottleneck function for concatenated features.r   )torchcatr)   r(   )r9   r>   concated_featuresbottleneck_outputs       r<   bottleneck_fnzDenseLayer.bottleneck_fn=   s2    !IIb!, JJtzz2C'DE  r=   xc                 .    |D ]  }|j                   s y y)z.Check if any tensor in list requires gradient.TF)requires_grad)r9   rE   tensors      r<   any_requires_gradzDenseLayer.any_requires_gradD   s"     	F##	 r=   c                 &      fd}t        |g| S )z5Call bottleneck function with gradient checkpointing.c                  &    j                  |       S )N)rD   )r>   r9   s    r<   closurez6DenseLayer.call_checkpoint_bottleneck.<locals>.closureM   s    %%b))r=   )r   )r9   rE   rL   s   `  r<   call_checkpoint_bottleneckz%DenseLayer.call_checkpoint_bottleneckK   s    	* '&A&&r=   c                    t        |t        j                        r|g}n|}| j                  rL| j	                  |      r;t        j
                  j                         rt        d      | j                  |      }n| j                  |      }| j                  | j                  |            }| j                  dkD  r,t        j                  || j                  | j                        }|S )zForward pass.

        Args:
            x: Input features (single tensor or list of tensors).

        Returns:
            New features to be concatenated.
        z%Memory Efficient not supported in JITr   )ptraining)
isinstancer@   Tensorr"   rI   jitis_scripting	ExceptionrM   rD   r/   r.   r!   FdropoutrP   )r9   rE   prev_featuresrC   new_featuress        r<   forwardzDenseLayer.forwardT   s     a&CMM""t'='=m'Lyy%%' GHH $ ? ? N $ 2 2= Azz$**->"?@>>A99\T^^dmm\Lr=   )__name__
__module____qualname____doc__r   intr   r6   Moduler8   boolr4   r
   r@   rR   rD   rI   rM   r	   rZ   __classcell__r;   s   @r<   r   r      s     +9!',5 #5 5 	5
 RYY5 5 !%5 
5@!U\\ 2 !u|| !4#5 $ 'D,> '5<< 'u||T%,,-??@ U\\ r=   r   c                        e Zd ZdZdZeddddfdededed	ed
eej                     de
deddf fdZdej                  dej                  fdZ xZS )
DenseBlockzTDenseNet Block.

    Contains multiple dense layers with concatenated features.
       r   FN
num_layersr   r   r   r    r!   r"   r#   c
           
          ||	d}
t         |           t        |      D ]2  }t        |||z  z   f|||||d|
}| j	                  d|dz   z  |       4 y)a  Initialize DenseBlock.

        Args:
            num_layers: Number of layers in the block.
            num_input_features: Number of input features.
            bn_size: Bottleneck size multiplier.
            growth_rate: Growth rate (k) for each layer.
            norm_layer: Normalization layer class.
            drop_rate: Dropout rate.
            grad_checkpointing: Use gradient checkpointing.
        r%   )r   r   r    r!   r"   zdenselayer%dr   N)r3   r4   ranger   r5   )r9   rg   r   r   r   r    r!   r"   r&   r'   r:   ilayerr;   s                r<   r4   zDenseBlock.__init__v   s{    . /z" 
	=A"Q_4'%##5 E OONa!e4e<
	=r=   init_featuresc                     |g}| j                         D ]  \  }} ||      }|j                  |         t        j                  |d      S )zForward pass through all layers in the block.

        Args:
            init_features: Initial features from previous layer.

        Returns:
            Concatenated features from all layers.
        r   )itemsappendr@   rA   )r9   rl   featuresnamerk   rY   s         r<   rZ   zDenseBlock.forward   sL     "?::< 	*KD% ?LOOL)	* yy1%%r=   )r[   r\   r]   r^   _versionr   r_   r   r6   r`   r8   ra   r4   r@   rR   rZ   rb   rc   s   @r<   re   re   o   s     H +9!',#=#= !$#= 	#=
 #= RYY#= #= !%#= 
#=J&U\\ &ell &r=   re   c                   |     e Zd ZdZedddfdededeej                     de	eej                        ddf
 fdZ
 xZS )	DenseTransitionzfTransition layer between DenseNet blocks.

    Reduces feature dimensions and spatial resolution.
    Nr   num_output_featuresr    aa_layerr#   c           
      >   ||d}t         |           | j                  d ||fi |       | j                  dt        j                  ||fdddd|       || j                  d ||fd	d
i|       y| j                  dt        j
                  d
d
             y)a  Initialize DenseTransition.

        Args:
            num_input_features: Number of input features.
            num_output_features: Number of output features.
            norm_layer: Normalization layer class.
            aa_layer: Anti-aliasing layer class.
        r%   normconvr   Fr*   Npoolr,   rf   )r+   r,   )r3   r4   r5   r6   r7   	AvgPool2d)	r9   r   ru   r    rv   r&   r'   r:   r;   s	           r<   r4   zDenseTransition.__init__   s    " /
+= D DE		 3!`AB1SX!`\^!` 	aOOFH-@$Q$Qb$QROOFBLLQq$IJr=   )r[   r\   r]   r^   r   r_   r   r6   r`   r   r4   rb   rc   s   @r<   rt   rt      sl     +926K #K "%K RYY	K
 tBII/K 
K Kr=   rt   c                   T    e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dde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dedededd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ddfdZdej.                  dej.                  fdZd dej.                  dedej.                  fdZdej.                  dej.                  fdZ xZS )#r   a  Densenet-BC model class.

    Based on `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`_

    Args:
        growth_rate: How many filters to add each layer (`k` in paper).
        block_config: How many layers in each pooling block.
        bn_size: Multiplicative factor for number of bottle neck layers
          (i.e. bn_size * k features in the bottleneck layer).
        drop_rate: Dropout rate before classifier layer.
        proj_drop_rate: Dropout rate after each dense layer.
        num_classes: Number of classification classes.
        memory_efficient: If True, uses checkpointing. Much more memory efficient,
          but slower. Default: *False*. See `"paper" <https://arxiv.org/pdf/1707.06990.pdf>`_.
    Nr   block_config.num_classesin_chansglobal_poolr   	stem_type	act_layerr    rv   r!   proj_drop_ratememory_efficientaa_stem_onlyr#   c                 X   ||d}|| _         || _        t        !|           t	        |	|      }	d|v }|dz  }|
t        j                  ddd      }n3t        j                  t        j                  ddd       |
d$|dd	|g }|r|x}}d
|v rd|dz  z  }d|v r|nd|dz  z  }t        j                  t        dt        j                  ||dfdddd|fd |	|fi |fdt        j                  ||dfdddd|fd |	|fi |fdt        j                  ||dfdddd|fd |	|fi |fd|fg            | _
        nMt        j                  t        dt        j                  ||fddddd|fd |	|fi |fd|fg            | _
        t        |dd|rdnd       g| _        d}|}t        |      D ]  \  }}t        d$|||||	||d|}d|dz    }| j                  j                  ||       |||z  z   }|rdn|
}|t!        |      dz
  k7  s]| xj                  t        ||d|z         gz  c_        |dz  }t#        d$||dz  |	|d|}| j                  j                  d |dz    |       |dz  } | j                  j                  d! |	|fi |       | xj                  t        ||d"      gz  c_        |x| _        | _        t)        | j$                  | j                   fd#|i|\  }}|| _        t        j,                  |      | _        || _        | j3                         D ]  } t5        | t
        j                        r*t
        j6                  j9                  | j:                         Gt5        | t
        j<                        rUt
        j6                  j?                  | j:                  d       t
        j6                  j?                  | j@                  d       t5        | t
        jB                        st
        j6                  j?                  | j@                  d        y)%aj  Initialize DenseNet.

        Args:
            growth_rate: How many filters to add each layer (k in paper).
            block_config: How many layers in each pooling block.
            num_classes: Number of classification classes.
            in_chans: Number of input channels.
            global_pool: Global pooling type.
            bn_size: Multiplicative factor for number of bottle neck layers.
            stem_type: Type of stem ('', 'deep', 'deep_tiered').
            act_layer: Activation layer.
            norm_layer: Normalization layer.
            aa_layer: Anti-aliasing layer.
            drop_rate: Dropout rate before classifier layer.
            proj_drop_rate: Dropout rate after each dense layer.
            memory_efficient: If True, uses checkpointing for memory efficiency.
            aa_stem_only: Apply anti-aliasing only to stem.
        r%   )r   deeprf   Nr0   r   )r+   r,   r2   )channelsr,   tiered   narrow   conv0F)r,   r2   r-   norm0r)   r(   r/   r.   pool0   r1   zfeatures.normr   )num_chs	reductionmodule)rg   r   r   r   r    r!   r"   
denseblockz	features.)r   ru   r    rv   
transitionnorm5zfeatures.norm5	pool_type )"r~   r   r3   r4   r   r6   	MaxPool2d
Sequentialr   r7   rp   dictfeature_info	enumeratere   r5   lenrt   num_featureshead_hidden_sizer   r   Dropout	head_drop
classifiermodulesrQ   initkaiming_normal_weightBatchNorm2d	constant_r-   Linear)"r9   r   r}   r~   r   r   r   r   r   r    rv   r!   r   r   r   r&   r'   r:   	deep_stemnum_init_features	stem_pool
stem_chs_1
stem_chs_2current_strider   rj   rg   blockmodule_nametransition_aa_layertransr   mr;   s"                                    r<   r4   zDenseNet.__init__   s   J /& '
iH
 i'	'!O1aHI1a@D"3ADD(F GI &11J9$+"23
2:i2G.QR]abRbMc
MM+"))Hj!cAqW\c`bcd*Z6267"))J
AeaQRY^ebdef*Z6267"))J0A1lQXY`eliklm*%6="=>)$7 + DM MM+"))H.?vQWXbcjovsuvw*%6="=>)$7 + DM *a-U^PQdeOf@ghj )&|4 	1MAz 	%#/'%(#3	 	E 'Ai0KMM$$[%8'*{*BBL*6$HC%))!!P[^iPij&l l!!#' '3(4(9)0	
  ((:a!eW)=uE+q07	1< 	  *\*HR*HId<>Zjkll4@@D1 #4#
 "#
 	#
Z 'I.$  	-A!RYY'''1Ar~~.!!!((A.!!!&&!,Aryy)!!!&&!,	-r=   coarsec                 8    t        d|rdn	ddt        fg      }|S )z"Group parameters for optimization.z<^features\.conv[012]|features\.norm[012]|features\.pool[012]z)^features\.(?:denseblock|transition)(\d+))z+^features\.denseblock(\d+)\.denselayer(\d+)Nz^features\.transition(\d+))stemblocks)r   r   )r9   r   matchers      r<   group_matcherzDenseNet.group_matcherd  s0     PCI?F.0@AP
 r=   enablec                 r    | j                   j                         D ]  }t        |t              s||_         y)z)Enable or disable gradient checkpointing.N)rp   r   rQ   r   r"   )r9   r   bs      r<   set_grad_checkpointingzDenseNet.set_grad_checkpointingp  s2     &&( 	.A!Z('-$	.r=   c                     | j                   S )zGet the classifier head.)r   )r9   s    r<   get_classifierzDenseNet.get_classifierw  s     r=   c                 p    || _         t        | 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   r   r   r   )r9   r~   r   s      r<   reset_classifierzDenseNet.reset_classifier|  s4     ',=t//;-H)$/r=   rE   c                 $    | j                  |      S )z/Forward pass through feature extraction layers.)rp   r9   rE   s     r<   forward_featureszDenseNet.forward_features  s    }}Qr=   
pre_logitsc                 p    | j                  |      }| j                  |      }|r|S | j                  |      S )zForward pass through classifier head.

        Args:
            x: Feature tensor.
            pre_logits: Return features before final classifier.

        Returns:
            Output tensor.
        )r   r   r   )r9   rE   r   s      r<   forward_headzDenseNet.forward_head  s8     QNN1q6DOOA$66r=   c                 J    | j                  |      }| j                  |      }|S )zoForward pass.

        Args:
            x: Input tensor.

        Returns:
            Output logits.
        )r   r   r   s     r<   rZ   zDenseNet.forward  s)     !!!$a r=   )    r              r0   avgr    relubatchnorm2dNr   r   FTNNF)T)r   )r[   r\   r]   r^   r_   r   strr   r   r6   r`   r8   ra   r4   r@   rS   ignorer   r   r   r   r   r   rR   r   r   rZ   rb   rc   s   @r<   r   r      s   $  ",;#$#+26!$&%*!%#D-D-  S/D- 	D-
 D- D- D- D- D- D- tBII/D- D- "D- #D- D-$ 
%D-L YY	D 	T#s(^ 	 	 YY.T .T . . YY		  	HC 	Hc 	Hd 	H %,,  5<<  7ell 7 7 7 %,, r=   
state_dictr#   c                     t        j                  d      }t        | j                               D ]D  }|j	                  |      }|s|j                  d      |j                  d      z   }| |   | |<   | |= F | S )zFilter torchvision pretrained state dict for compatibility.

    Args:
        state_dict: State dictionary from torchvision checkpoint.

    Returns:
        Filtered state dictionary.
    z]^(.*denselayer\d+\.(?:norm|relu|conv))\.((?:[12])\.(?:weight|bias|running_mean|running_var))$r   rf   )recompilelistkeysmatchgroup)r   patternkeyresnew_keys        r<   _filter_torchvision_pretrainedr     s{     jjhjG JOO%&  mmC iilSYYq\1G",S/Jw3  r=   variantr   r}   .
pretrainedc                 \    ||d<   ||d<   t        t        | |ft        d      t        d|S )a.  Create a DenseNet model.

    Args:
        variant: Model variant name.
        growth_rate: Growth rate parameter.
        block_config: Block configuration.
        pretrained: Load pretrained weights.
        **kwargs: Additional model arguments.

    Returns:
        DenseNet model instance.
    r   r}   T)flatten_sequential)feature_cfgpretrained_filter_fn)r   r   r   r   )r   r   r}   r   kwargss        r<   _create_densenetr     sJ    & (F=)F> D1;  r=   urlc                 2    | dddddt         t        dddd	|S )
z1Create default configuration for DenseNet models.r   )r0      r   )r   r   g      ?bicubiczfeatures.conv0r   z
apache-2.0)r   r~   
input_size	pool_sizecrop_pctinterpolationmeanstd
first_convr   licenser   )r   r   s     r<   _cfgr     s4     4}SYI%.B&l|	
  r=   ztimm/)r0      r   gffffff?)	hf_hub_idtest_input_sizetest_crop_pct)r   )zdensenet121.ra_in1kzdensenetblur121d.ra_in1kzdensenet264d.untraineddensenet121.tv_in1kzdensenet169.tv_in1kzdensenet201.tv_in1kzdensenet161.tv_in1kc           	      N    t        dd      }t        dd| it        |fi |}|S )ztDensenet-121 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    r   r   r   r}   r   )densenet121r   r   r   r   
model_argsmodels       r<   r  r    2    
 "?CJ`z`T*E_X^E_`ELr=   c           	      Z    t        dddt              }t        dd| it        |fi |}|S )zDensenet-121 w/ blur-pooling & 3-layer 3x3 stem
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    r   r   r   )r   r}   r   rv   r   )densenetblur121d)r   r   r   r  s       r<   r  r    s7    
 "?f_ijJeJe$zJd]cJdeELr=   c           	      N    t        dd      }t        dd| it        |fi |}|S )ztDensenet-169 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    r   )r   r   r   r   r   r   )densenet169r  r  s       r<   r
  r
  	  r  r=   c           	      N    t        dd      }t        dd| it        |fi |}|S )ztDensenet-201 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    r   )r   r   0   r   r   r   )densenet201r  r  s       r<   r  r    r  r=   c           	      N    t        dd      }t        dd| it        |fi |}|S )ztDensenet-161 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    r  )r   r   $   r   r   r   )densenet161r  r  s       r<   r  r    r  r=   c           	      P    t        ddd      }t        dd| it        |fi |}|S )ztDensenet-264 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    r  )r   r   @   r  r   )r   r}   r   r   )densenet264dr  r  s       r<   r  r  '  s4    
 "?fUJa
ad:F`Y_F`aELr=   tv_densenet121r   )r   r   )<r^   r   collectionsr   typingr   r   r   r   r   r	   r@   torch.nnr6   torch.nn.functional
functionalrV   torch.jit.annotationsr
   	timm.datar   r   timm.layersr   r   r   r   _builderr   _manipulater   r   	_registryr   r   r   __all__r`   r   
ModuleDictre   r   rt   r   r   r   rR   r   r_   ra   r   r   default_cfgsr  r  r
  r  r  r  r[   r   r=   r<   <module>r#     s   
 # : :     & A Y Y * 5 Y Y,U Up9& 9&xKbmm KDWryy Wtt S%,,=N8O * CHo 	 >c T#s(^  %%T; !%%T!; #f'2'2'2'2&  x   H   x   x   x      H+' r=   