a
    d.                     @   s@  d Z ddlmZ ddl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 ddlmZ dd	lmZ dd
lmZ ddlmZmZ dgZeeejdZG dd dejZd@ddZ dAddZ!G dd dejZ"dd Z#dBdd Z$ee$d!d"e$d!d"e$d!d"e$d!d"e$d!d"e$ e$ e$ e$d!d#dd$d%d&e$d!d#dd$d%d&e$d!d'd#dd$d%d(e$d!d'd#dd$d%d(d)Z%edCe"d+d,d-Z&edDe"d+d.d/Z'edEe"d+d0d1Z(edFe"d+d2d3Z)edGe"d+d4d5Z*edHe"d+d6d7Z+edIe"d+d8d9Z,edJe"d+d:d;Z-edKe"d+d<d=Z.edLe"d+d>d?Z/dS )Ma   ReXNet

A PyTorch impl of `ReXNet: Diminishing Representational Bottleneck on Convolutional Neural Network` -
https://arxiv.org/abs/2007.00992

Adapted from original impl at https://github.com/clovaai/rexnet
Copyright (c) 2020-present NAVER Corp. MIT license

Changes for timm, feature extraction, and rounded channel variant hacked together by Ross Wightman
Copyright 2020 Ross Wightman
    )partialceilNIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)ClassifierHeadcreate_act_layerConvNormActDropPathmake_divisibleSEModule   )build_model_with_cfg)efficientnet_init_weights)checkpoint_seq)generate_default_cfgsregister_modelRexNet)Z
norm_layerc                       s0   e Zd Zd fdd		ZdddZdd Z  ZS )LinearBottleneckr   r         ?        r   swishrelu6Nc              	      s   t t|   |dko,|d |d ko,||k| _|| _|| _|dkrjtt|| |d}t|||d| _	n
|}d | _	t||d||d |dd| _
|dkrt|tt|| |d	| _nd | _t|	| _t||ddd
| _|
| _d S )Nr   r   r   Zdivisor	act_layer   F)Zkernel_sizestridedilationgroups	apply_act)Zrd_channels)r"   )superr   __init__use_shortcutin_channelsout_channelsr   roundr
   conv_expconv_dw
SEWithNormintser	   act_dwconv_pwl	drop_path)selfin_chsout_chsr   r    	exp_ratiose_ratioch_divr   dw_act_layerr0   Zdw_chs	__class__ [/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/rexnet.pyr$   !   s0    "	
zLinearBottleneck.__init__Fc                 C   s   |r| j jS | jS N)r*   r'   )r1   expr:   r:   r;   feat_channelsL   s    zLinearBottleneck.feat_channelsc                 C   s   |}| j d ur|  |}| |}| jd ur6| |}| |}| |}| jr| jd urd| |}tj|d d d| j	f | |d d | j	d f gdd}|S )Nr   r   )Zdim)
r)   r*   r-   r.   r/   r%   r0   torchcatr&   )r1   xZshortcutr:   r:   r;   forwardO   s    








:zLinearBottleneck.forward)r   r   r   r   r   r   N)F)__name__
__module____qualname__r$   r>   rB   __classcell__r:   r:   r8   r;   r       s          +
r   r         r   c                    s  g dg d} fddD t fddt|D g }dgd  dgt dd    }t d d  d	 }| d
k r||  n|}	g }
t|d	 D ]2}|
tt|	|  |d |	||d	 d
  7 }	qdgd d   |gt dd    }tt|
|||S )N)r      rI   r   r      )r   rI   rI   rI   r   rI   c                    s   g | ]}t |  qS r:   r   ).0element)
depth_multr:   r;   
<listcomp>i       z_block_cfg.<locals>.<listcomp>c                    s(   g | ] \}}|gd g | d    qS )r   r:   )rK   idxrL   )layersr:   r;   rN   j   rO   r   r      r   r   r   r   rI   )sum	enumeraterangeappendr   r(   listzip)
width_multrM   initial_chs	final_chsr5   r6   stridesZ
exp_ratiosdepthZbase_chsZout_chs_listiZ	se_ratiosr:   )rM   rQ   r;   
_block_cfg_   s    $,r_       r   r   c                 C   sH  |g}g }	d}
d}g }t | }t| D ]\}\}}}}|}|dkr|dkrNdnd|d  }|	t|d |
|dg7 }	|
|kr|| }d}|| |d  }|dkrt|nd }|t||||||f|||||d	
 |
|9 }
|}|}||d  g7 }q&td
| |d}|	t|d |
dt |d  dg7 }	|t|||d ||	fS )NrI   r   r   stemz	features.)Znum_chsZ	reductionmoduler   )
r2   r3   r4   r   r    r5   r6   r   r7   r0   i   r   r   )	lenrT   dictr   rV   r   r>   r   r
   )	block_cfgZprev_chsrY   r6   output_strider   r7   drop_path_rateZfeat_chsfeature_infoZcurr_strider    featuresZ
num_blocksZ	block_idxZchsr4   r   r5   Znext_dilationfnameZ	block_dprr0   Zpen_chsr:   r:   r;   _build_blocksz   sH    
&rl   c                       s|   e Zd Zd! fdd	Zejjd"ddZejjd#ddZejjdd Z	d$ddZ
dd Zd%edddZdd  Z  ZS )&r   r     avgr`   rG   rH   r   UUUUUU?r   r   r   皙?r   c              	      s   t t|   || _|| _d| _|dv s,J |dk r<d| nd}tt|| |
d}t||dd|d| _	t
|||||	|
}t||||
||||\}| _|d	 j| _tj| | _t| j|||| _t|  d S )
NF)r`   rG      r   r`   r   r   rI   )r   r   rb   )r#   r   r$   num_classes	drop_rategrad_checkpointingr   r(   r
   ra   r_   rl   ri   r'   num_featuresnnZ
Sequentialrj   r   headr   )r1   Zin_chansrr   global_poolrg   rZ   r[   rY   rM   r5   r6   r   r7   rs   rh   Zstem_base_chsZstem_chsrf   rj   r8   r:   r;   r$      s.    

zRexNet.__init__Fc                 C   s   t ddd}|S )Nz^stemz^features\.(\d+))ra   blocks)re   )r1   ZcoarseZmatcherr:   r:   r;   group_matcher   s
    zRexNet.group_matcherTc                 C   s
   || _ d S r<   )rt   )r1   enabler:   r:   r;   set_grad_checkpointing   s    zRexNet.set_grad_checkpointingc                 C   s   | j jS r<   )rw   Zfc)r1   r:   r:   r;   get_classifier   s    zRexNet.get_classifierc                 C   s   t | j||| jd| _d S )N)Z	pool_typers   )r   ru   rs   rw   )r1   rr   rx   r:   r:   r;   reset_classifier   s    zRexNet.reset_classifierc                 C   s:   |  |}| jr,tj s,t| j|dd}n
| |}|S )NT)flatten)ra   rt   r?   jitZis_scriptingr   rj   r1   rA   r:   r:   r;   forward_features   s
    

zRexNet.forward_features
pre_logitsc                 C   s   |r| j ||dS |  |S )Nr   )rw   )r1   rA   r   r:   r:   r;   forward_head   s    zRexNet.forward_headc                 C   s   |  |}| |}|S r<   )r   r   r   r:   r:   r;   rB      s    

zRexNet.forward)r   rm   rn   r`   rG   rH   r   r   ro   r   r   r   rp   r   )F)T)rn   )F)rC   rD   rE   r$   r?   r   ignorerz   r|   r}   r~   r   boolr   rB   rF   r:   r:   r8   r;   r      s2                 -

c                 K   s"   t dd}tt| |fd|i|S )NT)Zflatten_sequentialfeature_cfg)re   r   r   )variant
pretrainedkwargsr   r:   r:   r;   _create_rexnet   s    
r    c                 K   s    | dddddt tdddd	|S )
Nrm   )r      r   )   r   g      ?Zbicubicz	stem.convzhead.fcZmit)urlrr   Z
input_sizeZ	pool_sizecrop_pctinterpolationmeanZstdZ
first_conv
classifierlicenser   )r   r   r:   r:   r;   _cfg  s    r   ztimm/)	hf_hub_idgffffff?)r      r   z
apache-2.0)r   r   test_crop_pcttest_input_sizer   i-.  )r   rr   r   r   r   r   )zrexnet_100.nav_in1kzrexnet_130.nav_in1kzrexnet_150.nav_in1kzrexnet_200.nav_in1kzrexnet_300.nav_in1kzrexnetr_100.untrainedzrexnetr_130.untrainedzrexnetr_150.untrainedzrexnetr_200.sw_in12k_ft_in1kzrexnetr_300.sw_in12k_ft_in1kzrexnetr_200.sw_in12kzrexnetr_300.sw_in12kF)returnc                 K   s   t d| fi |S )zReXNet V1 1.0x
rexnet_100r   r   r   r:   r:   r;   r   +  s    r   c                 K   s   t d| fddi|S )zReXNet V1 1.3x
rexnet_130rY   ?r   r   r:   r:   r;   r   1  s    r   c                 K   s   t d| fddi|S )zReXNet V1 1.5x
rexnet_150rY         ?r   r   r:   r:   r;   r   7  s    r   c                 K   s   t d| fddi|S )zReXNet V1 2.0x
rexnet_200rY          @r   r   r:   r:   r;   r   =  s    r   c                 K   s   t d| fddi|S )zReXNet V1 3.0x
rexnet_300rY         @r   r   r:   r:   r;   r   C  s    r   c                 K   s   t d| fddi|S )z*ReXNet V1 1.0x w/ rounded (mod 8) channelsrexnetr_100r6   rq   r   r   r:   r:   r;   r   I  s    r   c                 K   s   t d| fddd|S )z*ReXNet V1 1.3x w/ rounded (mod 8) channelsrexnetr_130r   rq   rY   r6   r   r   r:   r:   r;   r   O  s    r   c                 K   s   t d| fddd|S )z*ReXNet V1 1.5x w/ rounded (mod 8) channelsrexnetr_150r   rq   r   r   r   r:   r:   r;   r   U  s    r   c                 K   s   t d| fddd|S )z*ReXNet V1 2.0x w/ rounded (mod 8) channelsrexnetr_200r   rq   r   r   r   r:   r:   r;   r   [  s    r   c                 K   s   t d| fddd|S )z+ReXNet V1 3.0x w/ rounded (mod 16) channelsrexnetr_300r   rG   r   r   r   r:   r:   r;   r   a  s    r   )r   r   rG   rH   r   r   )r   r`   r   r   r   )r   )F)F)F)F)F)F)F)F)F)F)0__doc__	functoolsr   mathr   r?   Ztorch.nnrv   Z	timm.datar   r   Ztimm.layersr   r	   r
   r   r   r   Z_builderr   Z_efficientnet_builderr   Z_manipulater   	_registryr   r   __all__ZBatchNorm2dr+   Moduler   r_   rl   r   r   r   Zdefault_cfgsr   r   r   r   r   r   r   r   r   r   r:   r:   r:   r;   <module>   s    @      
     
0R

