
    ^jN                        d Z ddlZddlmZ ddl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mZ ddlmZ ddlmZ dd	lmZmZ d
gZd Z G d de
j4                        Z G d de
j4                        Z G d de      Z G d de      Z G d de      Z G d de
j4                        Z  G d d
e
j4                        Z!d2dZ"d3dZ# e e#d       e#dd       e#d        e#d!       e#d"       e#d#       e#d$d       e#d%       e#d&      d'	      Z$ed2d(e!fd)       Z%ed2d(e!fd*       Z&ed2d(e!fd+       Z'ed2d(e!fd,       Z(ed2d(e!fd-       Z)ed2d(e!fd.       Z*ed2d(e!fd/       Z+ed2d(e!fd0       Z,ed2d(e!fd1       Z-y)4a:  
SEResNet implementation from Cadene's pretrained models
https://github.com/Cadene/pretrained-models.pytorch/blob/master/pretrainedmodels/models/senet.py
Additional credit to https://github.com/creafz

Original model: https://github.com/hujie-frank/SENet

ResNet code gently borrowed from
https://github.com/pytorch/vision/blob/master/torchvision/models/resnet.py

FIXME I'm deprecating this model and moving them to ResNet as I don't want to maintain duplicate
support for extras like dilation, switchable BN/activations, feature extraction, etc that don't exist here.
    N)OrderedDict)TypeOptionalTupleIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)create_classifier   )build_model_with_cfg)register_modelgenerate_default_cfgsSENetc                 p   t        | t        j                        r-t        j                  j	                  | j
                  dd       y t        | t        j                        rUt        j                  j                  | j
                  d       t        j                  j                  | j                  d       y y )Nfan_outrelu)modenonlinearityg      ?        )	
isinstancennConv2dinitkaiming_normal_weightBatchNorm2d	constant_bias)ms    \/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/senet.py_weight_initr!      sp    !RYY
yvN	Ar~~	&
!((B'
!&&"% 
'    c                   0     e Zd Zddedef fdZd Z xZS )SEModulechannels	reductionc                    ||d}t         |           t        j                  |||z  fddi|| _        t        j
                  d      | _        t        j                  ||z  |fddi|| _        t        j                         | _	        y )Ndevicedtypekernel_sizer   Tinplace)
super__init__r   r   fc1ReLUr   fc2Sigmoidsigmoid)selfr%   r&   r)   r*   dd	__class__s         r    r/   zSEModule.__init__(   sw    /99Xx9'<R!RrRGGD)	99X2HR!RrRzz|r"   c                     |}|j                  dd      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }||z  S )N)      T)keepdim)meanr0   r   r2   r4   )r5   xmodule_inputs      r    forwardzSEModule.forward0   sX    FF64F(HHQKIIaLHHQKLLOar"   )NN)__name__
__module____qualname__intr/   r?   __classcell__r7   s   @r    r$   r$   &   s    $ $ $ r"   r$   c                       e Zd ZdZd Zy)
BottleneckzH
    Base class for bottlenecks that implements `forward()` method.
    c                    |}| j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }| j                  |      }| j                  |      }| j                  | j                  |      }| j                  |      |z   }| j                  |      }|S N)	conv1bn1r   conv2bn2conv3bn3
downsample	se_moduler5   r=   shortcutouts       r    r?   zBottleneck.forward?   s    jjmhhsmiinjjohhsmiinjjohhsm??&q)HnnS!H,iin
r"   N)r@   rA   rB   __doc__r?    r"   r    rG   rG   :   s    r"   rG   c                   d     e Zd ZdZdZ	 	 	 	 d
dedededededeej                     f fd	Z	 xZ
S )SEBottleneckz"
    Bottleneck for SENet154.
       inplanesplanesgroupsr&   striderP   c	           	      <   ||d}	t         
|           t        j                  ||dz  fddd|	| _        t        j
                  |dz  fi |	| _        t        j                  |dz  |dz  fd|d|dd|	| _        t        j
                  |dz  fi |	| _        t        j                  |dz  |dz  fddd|	| _	        t        j
                  |dz  fi |	| _
        t        j                  d	
      | _        t        |dz  fd|i|	| _        || _        || _        y )Nr(   r9   r   Fr+   r   rY   r:   r+   r]   paddingr\   r   Tr,   r&   r.   r/   r   r   rJ   r   rK   rL   rM   rN   rO   r1   r   r$   rQ   rP   r]   r5   rZ   r[   r\   r&   r]   rP   r)   r*   r6   r7   s             r    r/   zSEBottleneck.__init__\   s'    /YYx!UURTU
>>&1*33YYQJQJ	
 	
 	

 >>&1*33YYvz6A:W15WTVW
>>&1*33GGD)	!&1*H	HRH$r"   r   NNNr@   rA   rB   rU   	expansionrC   r   r   Moduler/   rD   rE   s   @r    rX   rX   V   sj     I .2  	
   !+ r"   rX   c                   d     e Zd ZdZdZ	 	 	 	 d
dedededededeej                     f fd	Z	 xZ
S )SEResNetBottleneckz
    ResNet bottleneck with a Squeeze-and-Excitation module. It follows Caffe
    implementation and uses `stride=stride` in `conv1` and not in `conv2`
    (the latter is used in the torchvision implementation of ResNet).
    rY   rZ   r[   r\   r&   r]   rP   c	                    ||d}	t         
|           t        j                  ||fdd|d|	| _        t        j
                  |fi |	| _        t        j                  ||fdd|dd|	| _        t        j
                  |fi |	| _        t        j                  ||dz  fddd|	| _	        t        j
                  |dz  fi |	| _
        t        j                  d	
      | _        t        |dz  fd|i|	| _        || _        || _        y )Nr(   r   Fr+   r   r]   r:   r+   ra   r\   r   rY   r_   Tr,   r&   rb   rc   s             r    r/   zSEResNetBottleneck.__init__   s     /YYx`QUSY`]_`
>>&/B/YYvvi1aPV]bifhi
>>&/B/YYvvzSquSPRS
>>&1*33GGD)	!&1*H	HRH$r"   rd   re   rE   s   @r    ri   ri   ~   sj    
 I .2  	
   !+ r"   ri   c                   j     e Zd ZdZdZ	 	 	 	 	 ddedededededeej                     d	ef fd
Z	 xZ
S )SEResNeXtBottleneckzI
    ResNeXt bottleneck type C with a Squeeze-and-Excitation module.
    rY   rZ   r[   r\   r&   r]   rP   
base_widthc
           	      V   ||	d}
t         |           t        j                  ||dz  z        |z  }t	        j
                  ||fdddd|
| _        t	        j                  |fi |
| _        t	        j
                  ||fd|d|dd|
| _	        t	        j                  |fi |
| _
        t	        j
                  ||dz  fddd	|
| _        t	        j                  |dz  fi |
| _        t	        j                  d
      | _        t        |dz  fd|i|
| _        || _        || _        y )Nr(   @   r   Frk   r:   r`   rY   r_   Tr,   r&   )r.   r/   mathfloorr   r   rJ   r   rK   rL   rM   rN   rO   r1   r   r$   rQ   rP   r]   )r5   rZ   r[   r\   r&   r]   rP   ro   r)   r*   r6   widthr7   s               r    r/   zSEResNeXtBottleneck.__init__   s    /

6Z"_56?YYxZAERSZWYZ
>>%.2.YYuev6ST]cjovsuv
>>%.2.YYufqjRaeRrR
>>&1*33GGD)	!&1*H	HRH$r"   )r   NrY   NNre   rE   s   @r    rn   rn      sw     I .2  	
   !+  r"   rn   c                   f     e Zd ZdZ	 	 	 	 d
dedededededeej                     f fdZd	 Z	 xZ
S )SEResNetBlockr   rZ   r[   r\   r&   r]   rP   c	                    ||d}	t         
|           t        j                  ||fdd|dd|	| _        t        j
                  |fi |	| _        t        j                  ||fdd|dd|	| _        t        j
                  |fi |	| _        t        j                  d      | _
        t        |fd	|i|	| _        || _        || _        y )
Nr(   r:   r   F)r+   ra   r]   r   rl   Tr,   r&   )r.   r/   r   r   rJ   r   rK   rL   rM   r1   r   r$   rQ   rP   r]   rc   s             r    r/   zSEResNetBlock.__init__   s     /YYxkQRX_dkhjk
>>&/B/YYvvi1aPV]bifhi
>>&/B/GGD)	!&DIDD$r"   c                 Z   |}| j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }| j
                  | j                  |      }| j                  |      |z   }| j                  |      }|S rI   )rJ   rK   r   rL   rM   rP   rQ   rR   s       r    r?   zSEResNetBlock.forward   s    jjmhhsmiinjjohhsmiin??&q)HnnS!H,iin
r"   rd   )r@   rA   rB   rf   rC   r   r   rg   r/   r?   rD   rE   s   @r    rv   rv      sc    I .2  	
   !+,r"   rv   c                       e Zd Z	 	 	 	 	 	 	 	 	 	 ddeej
                     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
f fdZ	 	 ddZej                  j                  dd       Zej                  j                  dd       Zej                  j                  dej
                  fd       Zddede
fdZd Zdde	fdZd Z xZS )r   blocklayers.r\   r&   	drop_ratein_chansrZ   	input_3x3downsample_kernel_sizedownsample_paddingnum_classesglobal_poolc                    t         |           ||d}|| _        || _        || _        || _        |rdt        j                  |ddfdddd|fd	t        j                  d)i |fd
t        j                  d      fdt        j                  ddddd|fdt        j                  d)i |fdt        j                  d      fdt        j                  d|dfdddd|fdt        j                  |fi |fdt        j                  d      fg	}nMdt        j                  ||fddddd|fd	t        j                  |fi |fd
t        j                  d      fg}t        j                  t        |            | _        t        j                  ddd      | _        t        |dd      g| _         | j"                  |fd|d   ||ddd|| _        | xj                   t        d|j&                  z  dd      gz  c_         | j"                  |fd|d   d|||	|
d|| _        | xj                   t        d|j&                  z  dd       gz  c_         | j"                  |fd!|d   d|||	|
d|| _        | xj                   t        d!|j&                  z  d"d#      gz  c_         | j"                  |fd$|d   d|||	|
d|| _        | xj                   t        d$|j&                  z  d%d&      gz  c_        d$|j&                  z  x| _        | _        t3        | j.                  | j                  fd'|i|\  | _        | _        | j9                         D ]  }t;        |        y()*af  
        Parameters
        ----------
        block (nn.Module): Bottleneck class.
            - For SENet154: SEBottleneck
            - For SE-ResNet models: SEResNetBottleneck
            - For SE-ResNeXt models:  SEResNeXtBottleneck
        layers (list of ints): Number of residual blocks for 4 layers of the
            network (layer1...layer4).
        groups (int): Number of groups for the 3x3 convolution in each
            bottleneck block.
            - For SENet154: 64
            - For SE-ResNet models: 1
            - For SE-ResNeXt models:  32
        reduction (int): Reduction ratio for Squeeze-and-Excitation modules.
            - For all models: 16
        dropout_p (float or None): Drop probability for the Dropout layer.
            If `None` the Dropout layer is not used.
            - For SENet154: 0.2
            - For SE-ResNet models: None
            - For SE-ResNeXt models: None
        inplanes (int):  Number of input channels for layer1.
            - For SENet154: 128
            - For SE-ResNet models: 64
            - For SE-ResNeXt models: 64
        input_3x3 (bool): If `True`, use three 3x3 convolutions instead of
            a single 7x7 convolution in layer0.
            - For SENet154: True
            - For SE-ResNet models: False
            - For SE-ResNeXt models: False
        downsample_kernel_size (int): Kernel size for downsampling convolutions
            in layer2, layer3 and layer4.
            - For SENet154: 3
            - For SE-ResNet models: 1
            - For SE-ResNeXt models: 1
        downsample_padding (int): Padding for downsampling convolutions in
            layer2, layer3 and layer4.
            - For SENet154: 1
            - For SE-ResNet models: 0
            - For SE-ResNeXt models: 0
        num_classes (int): Number of outputs in `last_linear` layer.
            - For all models: 1000
        r(   rJ   rq   r:   r9   r   F)r]   ra   r   rK   relu1Tr,   rL   )rq   rq   r:   rM   relu2rN   rO   relu3   r+   r]   ra   r   )r]   	ceil_modelayer0)num_chsr&   moduler   )r[   blocksr\   r&   r   r   rY   layer1   )r[   r   r]   r\   r&   r   r      layer2      layer3i       layer4	pool_typeN)rq   )r.   r/   rZ   r   r}   r|   r   r   r   r1   
Sequentialr   r   	MaxPool2dpool0dictfeature_info_make_layerr   rf   r   r   r   num_featureshead_hidden_sizer
   r   last_linearmodulesr!   )r5   rz   r{   r\   r&   r|   r}   rZ   r~   r   r   r   r   r)   r*   r6   layer0_modulesr   r7   s                     r    r/   zSENet.__init__   s   x 	/ & ""))Hb![Aqu[XZ[\0R01"''$/0"))UaURTUV0R01"''$/0"))B![Aqu[XZ[\x6267"''$/0
N "))HhmAaYZafmjlmnx6267"''$/0N
 mmK$?@\\!A>
!(aQR&d&&	
!9#$ 	
 	
 	d2+?1U]^__&d&&

!9#91

 

 	d3+@AV^_``&d&&

!9#91

 

 	d3+@BW_`aa&d&&

!9#91

 

 	d3+@BW_`aa47%//4IID1->.
 ".
 	.
*$*  	AO	r"   c           
         |	|
d}d }|dk7  s| j                   ||j                  z  k7  rht        j                  t        j                  | j                   ||j                  z  f|||dd|t        j
                  ||j                  z  fi |      } || j                   |||||fi |g}||j                  z  | _         t        d|      D ]'  }|j                   || j                   |||fi |       ) t        j                  | S )Nr(   r   Fr   )rZ   rf   r   r   r   r   rangeappend)r5   rz   r[   r   r\   r&   r]   r   r   r)   r*   r6   rP   r{   is                  r    r   zSENet._make_layer  s   /
Q;$--6EOO+CC		MM6EOO#;QI_!+=EQMOQ v7>2>	J vvy&*[XZ[\0q&! 	QAMM%vvyOBOP	Q }}f%%r"   c                 (    t        d|rdnd      }|S )Nz^layer0z^layer(\d+)z^layer(\d+)\.(\d+))stemr   )r   )r5   coarsematchers      r    group_matcherzSENet.group_matcher  s    J~Mbcr"   c                     |rJ d       y )Nz$gradient checkpointing not supportedrV   )r5   enables     r    set_grad_checkpointingzSENet.set_grad_checkpointing  s    AAAz6r"   returnc                     | j                   S rI   )r   )r5   s    r    get_classifierzSENet.get_classifier  s    r"   c                 p    || _         t        | j                  | j                   |      \  | _        | _        y )N)r   )r   r
   r   r   r   )r5   r   r   s      r    reset_classifierzSENet.reset_classifier  s3    &->t//;.H*$*r"   c                     | j                  |      }| j                  |      }| j                  |      }| j                  |      }| j	                  |      }| j                  |      }|S rI   )r   r   r   r   r   r   r5   r=   s     r    forward_featureszSENet.forward_features  sU    KKNJJqMKKNKKNKKNKKNr"   
pre_logitsc                     | j                  |      }| j                  dkD  r,t        j                  || j                  | j                        }|r|S | j                  |      S )Nr   )ptraining)r   r|   Fdropoutr   r   )r5   r=   r   s      r    forward_headzSENet.forward_head  sP    Q>>B		!t~~FAq7D$4$4Q$77r"   c                 J    | j                  |      }| j                  |      }|S rI   )r   r   r   s     r    r?   zSENet.forward  s'    !!!$a r"   )
g?r:   rq   Fr   r     avgNN)r   r   r   NNF)T)r   )r@   rA   rB   r   r   rg   r   rC   floatboolstrr/   r   torchjitignorer   r   r   r   r   r   r?   rD   rE   s   @r    r   r      s^     ##*+&'#$P		?P #s(OP 	P
 P P P P P %(P !$P P Pd LMW[&& YY  YYB B YY 		    HC Hc H
8$ 8r"   c                 &    t        t        | |fi |S rI   )r   r   )variant
pretrainedkwargss      r    _create_senetr     s    w
EfEEr"   c                 2    | dddddt         t        dddd	|S )
Nr   )r:      r   )r   r   g      ?bilinearzlayer0.conv1r   z
apache-2.0)urlr   
input_size	pool_sizecrop_pctinterpolationr<   std
first_conv
classifierlicenser   )r   r   s     r    _cfgr     s2    4}SYJ%.B$Ml	
  r"   zmhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/legacy_senet154-e9eb9fe6.pth)r   zhhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/seresnet18-4bb0ce65.pthbicubic)r   r   zhhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/seresnet34-a4004e63.pthzhhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-cadene/se_resnet50-ce0d4300.pthzihttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-cadene/se_resnet101-7e38fcc6.pthzihttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-cadene/se_resnet152-d17c99b7.pthzphttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/seresnext26_32x4d-65ebdb501.pthzwhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/legacy_se_resnext50_32x4d-f3651bad.pthzxhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/legacy_se_resnext101_32x4d-37725eac.pth)	zlegacy_senet154.in1kzlegacy_seresnet18.in1kzlegacy_seresnet34.in1kzlegacy_seresnet50.in1kzlegacy_seresnet101.in1kzlegacy_seresnet152.in1kzlegacy_seresnext26_32x4d.in1kzlegacy_seresnext50_32x4d.in1kzlegacy_seresnext101_32x4d.in1kr   c           	      Z    t        t        g ddd      }t        d| fi t        |fi |S )Nr9   r9   r9   r9   r   r   rz   r{   r\   r&   legacy_seresnet18r   rv   r   r   r   
model_argss      r    r   r     3    LbJJ,jWD<Vv<VWWr"   c           	      Z    t        t        g ddd      }t        d| fi t        |fi |S )Nr:   rY      r:   r   r   r   legacy_seresnet34r   r   s      r    r   r     r   r"   c           	      Z    t        t        g ddd      }t        d| fi t        |fi |S )Nr   r   r   r   legacy_seresnet50r   ri   r   r   s      r    r   r     s3     a2OJ,jWD<Vv<VWWr"   c           	      Z    t        t        g ddd      }t        d| fi t        |fi |S )Nr:   rY      r:   r   r   r   legacy_seresnet101r   r   s      r    r   r     4     qBPJ-zXT*=WPV=WXXr"   c           	      Z    t        t        g ddd      }t        d| fi t        |fi |S )Nr:   r   $   r:   r   r   r   legacy_seresnet152r   r   s      r    r   r     r   r"   c           
      b    t        t        g ddddddd      }t        d	| fi t        |fi |S )
Nr   rq   r   r:   r   r   T)rz   r{   r\   r&   r   r   rZ   r~   legacy_senet154)r   rX   r   r   s      r    r   r     s?    =r Q#QUWJ *JU$z:TV:TUUr"   c           	      Z    t        t        g ddd      }t        d| fi t        |fi |S )Nr   r   r   r   legacy_seresnext26_32x4dr   rn   r   r   s      r    r   r     4    !,rRQJ3Z^4
C]V\C]^^r"   c           	      Z    t        t        g ddd      }t        d| fi t        |fi |S )Nr   r   r   r   legacy_seresnext50_32x4dr   r   s      r    r   r     r   r"   c           	      Z    t        t        g ddd      }t        d| fi t        |fi |S )Nr   r   r   r   legacy_seresnext101_32x4dr   r   s      r    r   r     s4    !-bRJ4j_DD^W]D^__r"   r   ) ).rU   rr   collectionsr   typingr   r   r   r   torch.nnr   torch.nn.functional
functionalr   	timm.datar   r	   timm.layersr
   _builderr   	_registryr   r   __all__r!   rg   r$   rG   rX   ri   rn   rv   r   r   r   default_cfgsr   r   r   r   r   r   r   r   r   rV   r"   r    <module>r	     sK    # ( (     A ) * <)& ryy  ( 8%: %P B* B*BII *ZKBII K\F % {}"v! #vx"vx#w y#w y%)~&! &* F&G&* G'H'& 0 XU X X XU X X XU X X Ye Y Y Ye Y Y V5 V V _E _ _ _E _ _ `U ` `r"   