a
    d                     @   s  d Z ddlmZmZmZ ddlZddlmZ ddlm  m	Z
 ddlmZmZ eeeeef f ZdddZded	d
dZd ed	ddZd!ed	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jZG dd dejZG dd dejZdS )"a   PyTorch selectable adaptive pooling
Adaptive pooling with the ability to select the type of pooling from:
    * 'avg' - Average pooling
    * 'max' - Max pooling
    * 'avgmax' - Sum of average and max pooling re-scaled by 0.5
    * 'avgmaxc' - Concatenation of average and max pooling along feature dim, doubles feature dim

Both a functional and a nn.Module version of the pooling is provided.

Hacked together by / Copyright 2020 Ross Wightman
    )OptionalTupleUnionN   )get_spatial_dimget_channel_dimavgc                 C   s   |  drdS dS d S )N	catavgmax   r   )endswith	pool_type r   i/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/layers/adaptive_avgmax_pool.pyadaptive_pool_feat_mult   s    
r   output_sizec                 C   s$   t | |}t | |}d||  S )N      ?)Fadaptive_avg_pool2dadaptive_max_pool2dxr   x_avgx_maxr   r   r   adaptive_avgmax_pool2d   s    r   c                 C   s(   t | |}t | |}t||fdS Nr   )r   r   r   torchcatr   r   r   r   adaptive_catavgmax_pool2d$   s    r   c                 C   sh   |dkrt | |} nN|dkr*t| |} n:|dkr>t| |} n&|dkrTt | |} ndsdJ d| | S )zFSelectable global pooling function with dynamic input kernel size
    r   avgmaxr	   maxFzInvalid pool type: %s)r   r   r   r   r   )r   r   r   r   r   r   select_adaptive_pool2d*   s    r"   c                       s.   e Zd Zdeed fddZdd Z  ZS )	FastAdaptiveAvgPoolFNCHWflatten	input_fmtc                    s"   t t|   || _t|| _d S N)superr#   __init__r&   r   dimselfr&   r'   	__class__r   r   r*   ;   s    zFastAdaptiveAvgPool.__init__c                 C   s   |j | j| j dS NZkeepdim)meanr+   r&   r-   r   r   r   r   forward@   s    zFastAdaptiveAvgPool.forward)Fr$   )__name__
__module____qualname__boolr   r*   r4   __classcell__r   r   r.   r   r#   :   s   r#   c                       s.   e Zd Zdeed fddZdd Z  ZS )	FastAdaptiveMaxPoolFr$   r%   c                    s"   t t|   || _t|| _d S r(   )r)   r:   r*   r&   r   r+   r,   r.   r   r   r*   E   s    zFastAdaptiveMaxPool.__init__c                 C   s   |j | j| j dS r0   )amaxr+   r&   r3   r   r   r   r4   J   s    zFastAdaptiveMaxPool.forward)Fr$   r5   r6   r7   r8   strr*   r4   r9   r   r   r.   r   r:   D   s   r:   c                       s.   e Zd Zdeed fddZdd Z  ZS )	FastAdaptiveAvgMaxPoolFr$   r%   c                    s"   t t|   || _t|| _d S r(   )r)   r>   r*   r&   r   r+   r,   r.   r   r   r*   O   s    zFastAdaptiveAvgMaxPool.__init__c                 C   s8   |j | j| j d}|j| j| j d}d| d|  S )Nr1   r   )r2   r+   r&   r;   r-   r   r   r   r   r   r   r4   T   s    zFastAdaptiveAvgMaxPool.forward)Fr$   r<   r   r   r.   r   r>   N   s   r>   c                       s.   e Zd Zdeed fddZdd Z  ZS )	FastAdaptiveCatAvgMaxPoolFr$   r%   c                    s8   t t|   || _t|| _|r*d| _n
t|| _d S r   )r)   r@   r*   r&   r   
dim_reducedim_catr   r,   r.   r   r   r*   [   s    
z"FastAdaptiveCatAvgMaxPool.__init__c                 C   s:   |j | j| j d}|j| j| j d}t||f| jS r0   )r2   rA   r&   r;   r   r   rB   r?   r   r   r   r4   d   s    z!FastAdaptiveCatAvgMaxPool.forward)Fr$   r<   r   r   r.   r   r@   Z   s   	r@   c                       s,   e Zd Zded fddZdd Z  ZS )AdaptiveAvgMaxPool2dr   r   c                    s   t t|   || _d S r(   )r)   rC   r*   r   r-   r   r.   r   r   r*   k   s    zAdaptiveAvgMaxPool2d.__init__c                 C   s   t || jS r(   )r   r   r3   r   r   r   r4   o   s    zAdaptiveAvgMaxPool2d.forward)r   r5   r6   r7   _int_tuple_2_tr*   r4   r9   r   r   r.   r   rC   j   s   rC   c                       s,   e Zd Zded fddZdd Z  ZS )AdaptiveCatAvgMaxPool2dr   r   c                    s   t t|   || _d S r(   )r)   rG   r*   r   rD   r.   r   r   r*   t   s    z AdaptiveCatAvgMaxPool2d.__init__c                 C   s   t || jS r(   )r   r   r3   r   r   r   r4   x   s    zAdaptiveCatAvgMaxPool2d.forward)r   rE   r   r   r.   r   rG   s   s   rG   c                       sN   e Zd ZdZdeeeed fddZd	d
 Zdd Z	dd Z
dd Z  ZS )SelectAdaptivePool2dzCSelectable global pooling layer with dynamic input kernel size
    r   fastFr$   )r   r   r&   r'   c                    sN  t t|   |dv sJ |p d| _|sLt | _|r@tdnt | _n|	ds^|dkr|dksnJ d|
drt||d| _nB|
d	rt||d| _n(|
d
rt||d| _nt||d| _t | _nt|dksJ |dkrt|| _n:|d	krt|| _n$|d
kr$t|| _nt|| _|r@tdnt | _d S )N)r$   ZNHWC r   rI   r$   zAFast pooling and non NCHW input formats require output_size == 1.r    )r'   r	   r!   )r)   rH   r*   r   nnZIdentitypoolZFlattenr&   
startswithr   r>   r@   r:   r#   rC   rG   ZAdaptiveMaxPool2dZAdaptiveAvgPool2d)r-   r   r   r&   r'   r.   r   r   r*      s2    






zSelectAdaptivePool2d.__init__c                 C   s   | j  S r(   r   r-   r   r   r   is_identity   s    z SelectAdaptivePool2d.is_identityc                 C   s   |  |}| |}|S r(   )rL   r&   r3   r   r   r   r4      s    

zSelectAdaptivePool2d.forwardc                 C   s
   t | jS r(   )r   r   rN   r   r   r   	feat_mult   s    zSelectAdaptivePool2d.feat_multc                 C   s(   | j jd d | j d t| j d S )Nz (z
pool_type=z
, flatten=))r/   r5   r   r=   r&   rN   r   r   r   __repr__   s    
zSelectAdaptivePool2d.__repr__)r   rI   Fr$   )r5   r6   r7   __doc__rF   r=   r8   r*   rO   r4   rP   rR   r9   r   r   r.   r   rH   |   s       $rH   )r   )r   )r   )r   r   )rS   typingr   r   r   r   Ztorch.nnrK   Ztorch.nn.functionalZ
functionalr   formatr   r   intrF   r   r   r   r"   Moduler#   r:   r>   r@   rC   rG   rH   r   r   r   r   <module>   s"   


		