a
    d                     @   s^   d Z ddlZddlm  mZ ddlmZ ddlmZ G dd dejZ	G dd	 d	ejZ
dS )
a'   Split Attention Conv2d (for ResNeSt Models)

Paper: `ResNeSt: Split-Attention Networks` - /https://arxiv.org/abs/2004.08955

Adapted from original PyTorch impl at https://github.com/zhanghang1989/ResNeSt

Modified for torchscript compat, performance, and consistency with timm by Ross Wightman
    N)nn   )make_divisiblec                       s$   e Zd Z fddZdd Z  ZS )RadixSoftmaxc                    s   t t|   || _|| _d S )N)superr   __init__radixcardinality)selfr   r	   	__class__ _/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/layers/split_attn.pyr      s    zRadixSoftmax.__init__c                 C   sZ   | d}| jdkrL||| j| jddd}tj|dd}||d}n
t	|}|S )Nr   r      Zdim)
sizer   viewr	   Z	transposeFZsoftmaxreshapetorchZsigmoid)r
   xbatchr   r   r   forward   s    


zRadixSoftmax.forward)__name__
__module____qualname__r   r   __classcell__r   r   r   r   r      s   r   c                       sH   e Zd ZdZdddddddddddejddf fd	d
	Zdd Z  ZS )	SplitAttnz Split-Attention (aka Splat)
    N   r   Fr   g      ?   c                    s  t t|   |p|}|	| _||	 }|d u rDt||	 |
 d|d}n||	 }|d u r\|d n|}tj||||||f||	 |d|| _|r||nt | _	|d ur| nt | _
|dd| _tj||d|d| _|r||nt | _|dd| _tj||d|d| _t|	|| _d S )	N    )Z	min_valueZdivisorr   )groupsbiasT)Zinplacer   )r"   )r   r   r   r   r   r   ZConv2dconvZIdentitybn0dropact0fc1bn1act1fc2r   rsoftmax)r
   Zin_channelsZout_channelsZkernel_sizeZstridepaddingZdilationr"   r#   r   Zrd_ratioZrd_channelsZ
rd_divisorZ	act_layerZ
norm_layerZ
drop_layerkwargsZmid_chsZattn_chsr   r   r   r   $   s.    zSplitAttn.__init__c           	      C   s   |  |}| |}| |}| |}|j\}}}}| jdkrj||| j|| j ||f}|jdd}n|}|jddd}| 	|}| 
|}| |}| |}| ||ddd}| jdkr|||| j|| j ddf jdd}n|| }| S )Nr   r   )r   r   T)Zkeepdimr   )r$   r%   r&   r'   shaper   r   summeanr(   r)   r*   r+   r,   r   
contiguous)	r
   r   BRCHWZx_gapZx_attnoutr   r   r   r   =   s&    









*zSplitAttn.forward)	r   r   r   __doc__r   ZReLUr   r   r   r   r   r   r   r   !   s   r   )r8   r   Ztorch.nn.functionalr   Z
functionalr   Zhelpersr   Moduler   r   r   r   r   r   <module>   s   