a
    dt/                     @   s   d Z ddlZddlmZ ddlmZ ddlmZmZm	Z	m
Z
mZ g dZdd ZG dd	 d	ejZG d
d dejZG dd dejZG dd dejZG dd deZG dd dejZdS )z[ EfficientNet, MobileNetV3, etc Blocks

Hacked together by / Copyright 2019, Ross Wightman
    N)
functional)create_conv2dDropPathmake_divisiblecreate_act_layerget_norm_act_layer)SqueezeExcite	ConvBnActDepthwiseSeparableConvInvertedResidualCondConvResidualEdgeResidualc                 C   s$   | sdS ||  dksJ ||  S d S )N   r    )
group_sizeZchannelsr   r   i/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/_efficientnet_blocks.py
num_groups   s    r   c                       s:   e Zd ZdZddejejddf fdd	Zdd Z  Z	S )r   a   Squeeze-and-Excitation w/ specific features for EfficientNet/MobileNet family

    Args:
        in_chs (int): input channels to layer
        rd_ratio (float): ratio of squeeze reduction
        act_layer (nn.Module): activation layer of containing block
        gate_layer (Callable): attention gate function
        force_act_layer (nn.Module): override block's activation fn if this is set/bound
        rd_round_fn (Callable): specify a fn to calculate rounding of reduced chs
    g      ?Nc                    sv   t t|   |d u r*|pt}||| }|p0|}tj||ddd| _t|dd| _tj||ddd| _	t|| _
d S )Nr   T)Zbiasinplace)superr   __init__roundnnZConv2dconv_reducer   act1conv_expandgate)selfin_chsZrd_ratioZrd_channels	act_layerZ
gate_layerZforce_act_layerZrd_round_fn	__class__r   r   r   %   s    zSqueezeExcite.__init__c                 C   s:   |j ddd}| |}| |}| |}|| | S )N)      T)Zkeepdim)meanr   r   r   r   )r   xZx_ser   r   r   forward2   s
    


zSqueezeExcite.forward)
__name__
__module____qualname____doc__r   ReLUZSigmoidr   r&   __classcell__r   r   r    r   r      s
   r   c                       sF   e Zd ZdZdddddejejdf fdd	Zd	d
 Zdd Z	  Z
S )r	   z@ Conv + Norm Layer + Activation w/ optional skip connection
    r   r    F        c              	      sx   t t|   t|
|	}t||}|o4|dko4||k| _t|||||||d| _||dd| _|rjt	|nt
 | _d S )Nr   stridedilationgroupspaddingTr   )r   r	   r   r   r   has_skipr   convbn1r   r   Identity	drop_path)r   r   out_chsZkernel_sizer0   r1   r   pad_typeskipr   
norm_layerdrop_path_ratenorm_act_layerr2   r    r   r   r   =   s    

zConvBnAct.__init__c                 C   s0   |dkrt dd| jjdS t dd| jjdS d S )N	expansionr6   r&   moduleZ	hook_typeZnum_chsr-   )dictr5   out_channelsr   locationr   r   r   feature_infoJ   s    zConvBnAct.feature_infoc                 C   s0   |}|  |}| |}| jr,| || }|S N)r5   r6   r4   r8   r   r%   shortcutr   r   r   r&   P   s    

zConvBnAct.forwardr'   r(   r)   r*   r   r+   BatchNorm2dr   rF   r&   r,   r   r   r    r   r	   :   s   r	   c                       sN   e Zd ZdZddddddddejejddf fdd		Zd
d Zdd Z	  Z
S )r
   z DepthwiseSeparable block
    Used for DS convs in MobileNet-V1 and in the place of IR blocks that have no expansion
    (factor of 1.0). This is an alternative to having a IR with an optional first pw conv.
    r#   r   r-   FNr.   c              	      s   t t|   t||}t||}|dko6||ko6| | _|
| _t|||||||d| _||dd| _	|rv|||dnt
 | _t|||	|d| _||d| jd| _|rt|nt
 | _d S )Nr   )r0   r1   r3   r2   Tr   r   r3   )r   	apply_act)r   r
   r   r   r   r4   Z
has_pw_actr   conv_dwr6   r   r7   seconv_pwbn2r   r8   )r   r   r9   dw_kernel_sizer0   r1   r   r:   noskippw_kernel_sizeZpw_actr   r<   se_layerr=   r>   r2   r    r   r   r   ^   s    

zDepthwiseSeparableConv.__init__c                 C   s0   |dkrt dd| jjdS t dd| jjdS d S )Nr?   rQ   forward_prer@   r-   )rB   rQ   in_channelsrC   rD   r   r   r   rF   s   s    z#DepthwiseSeparableConv.feature_infoc                 C   sN   |}|  |}| |}| |}| |}| |}| jrJ| || }|S rG   )rO   r6   rP   rQ   rR   r4   r8   rH   r   r   r   r&   y   s    




zDepthwiseSeparableConv.forwardrJ   r   r   r    r   r
   Y   s   
r
   c                       sR   e Zd ZdZdddddddddejejdddf fd	d
	Zdd Zdd Z	  Z
S )r   a   Inverted residual block w/ optional SE

    Originally used in MobileNet-V2 - https://arxiv.org/abs/1801.04381v4, this layer is often
    referred to as 'MBConv' for (Mobile inverted bottleneck conv) and is also used in
      * MNasNet - https://arxiv.org/abs/1807.11626
      * EfficientNet - https://arxiv.org/abs/1905.11946
      * MobileNet-V3 - https://arxiv.org/abs/1905.02244
    r#   r   r-   F      ?Nr.   c                    s   t t|   t||}|pi }t||	 }t||}||koJ|dkoJ| | _t|||
fd|i|| _||dd| _	t|||f||||d|| _
||dd| _|r|||dnt | _t|||fd|i|| _||dd| _|rt|nt | _d S )	Nr   r3   Tr   r/   rL   FrN   )r   r   r   r   r   r   r4   r   rQ   r6   rO   rR   r   r7   rP   conv_pwlbn3r   r8   )r   r   r9   rS   r0   r1   r   r:   rT   	exp_ratioexp_kernel_sizerU   r   r<   rV   conv_kwargsr=   r>   mid_chsr2   r    r   r   r      s*    

zInvertedResidual.__init__c                 C   s0   |dkrt dd| jjdS t dd| jjdS d S Nr?   r[   rW   r@   r-   rB   r[   rX   rC   rD   r   r   r   rF      s    zInvertedResidual.feature_infoc                 C   sb   |}|  |}| |}| |}| |}| |}| |}| |}| jr^| || }|S rG   )	rQ   r6   rO   rR   rP   r[   r\   r4   r8   rH   r   r   r   r&      s    






zInvertedResidual.forwardrJ   r   r   r    r   r      s   


r   c                       sJ   e Zd ZdZdddddddddejejddd	f fd
d	Zdd Z  Z	S )r   z, Inverted residual block w/ CondConv routingr#   r   r-   FrY   Nr   r.   c                    sV   || _ t| j d}tt| j||||||||||	|
|||||d t|| j | _d S )N)num_experts)rS   r0   r1   r   r:   r   rT   r]   r^   rU   rV   r<   r_   r=   )rc   rB   r   r   r   r   ZLinear
routing_fn)r   r   r9   rS   r0   r1   r   r:   rT   r]   r^   rU   r   r<   rV   rc   r=   r_   r    r   r   r      s    

zCondConvResidual.__init__c                 C   s   |}t |dd}t| |}| ||}| |}| ||}| 	|}| 
|}| ||}| |}| jr| || }|S )Nr   )FZadaptive_avg_pool2dflattentorchZsigmoidrd   rQ   r6   rO   rR   rP   r[   r\   r4   r8   )r   r%   rI   Zpooled_inputsZrouting_weightsr   r   r   r&      s    



zCondConvResidual.forward)
r'   r(   r)   r*   r   r+   rK   r   r&   r,   r   r   r    r   r      s   

r   c                       sP   e Zd ZdZdddddddddejejdd	f fd
d	Zdd Zdd Z	  Z
S )r   a(   Residual block with expansion convolution followed by pointwise-linear w/ stride

    Originally introduced in `EfficientNet-EdgeTPU: Creating Accelerator-Optimized Neural Networks with AutoML`
        - https://ai.googleblog.com/2019/08/efficientnet-edgetpu-creating.html

    This layer is also called FusedMBConv in the MobileDet, EfficientNet-X, and EfficientNet-V2 papers
      * MobileDet - https://arxiv.org/abs/2004.14525
      * EfficientNet-X - https://arxiv.org/abs/2102.05610
      * EfficientNet-V2 - https://arxiv.org/abs/2104.00298
    r#   r   r   r-   FrY   Nr.   c              	      s   t t|   t||}|dkr.t||
 }nt||
 }t||}||koX|dkoX|	 | _t|||||||d| _||dd| _	|r|||dnt
 | _t||||d| _||dd	| _|rt|nt
 | _d S )
Nr   r   r/   Tr   rL   rM   FrZ   )r   r   r   r   r   r   r4   r   conv_expr6   r   r7   rP   r[   rR   r   r8   )r   r   r9   r^   r0   r1   r   r:   Zforce_in_chsrT   r]   rU   r   r<   rV   r=   r>   r`   r2   r    r   r   r      s    

zEdgeResidual.__init__c                 C   s0   |dkrt dd| jjdS t dd| jjdS d S ra   rb   rD   r   r   r   rF   	  s    zEdgeResidual.feature_infoc                 C   sN   |}|  |}| |}| |}| |}| |}| jrJ| || }|S rG   )rh   r6   rP   r[   rR   r4   r8   rH   r   r   r   r&     s    




zEdgeResidual.forwardrJ   r   r   r    r   r      s   
r   )r*   rg   Ztorch.nnr   r   re   Ztimm.layersr   r   r   r   r   __all__r   Moduler   r	   r
   r   r   r   r   r   r   r   <module>   s   	!,;#