a
    d>                     @   s   d Z ddlm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Z ddlmZ dd	lmZ dd
lmZ deeeeee dddZdddZdeeeeeee dddZG dd dejZG dd dejZdS )zV Classifier head and layer factory

Hacked together by / Copyright 2020 Ross Wightman
    )OrderedDict)partial)OptionalUnionCallableN)
functional   )SelectAdaptivePool2d)get_act_layer)get_norm_layeravgF)num_featuresnum_classes	pool_typeuse_conv	input_fmtc                 C   sD   | }|s"|dks|sJ dd}t |||d}| |  }||fS )Nr   zUPooling can only be disabled if classifier is also removed or conv classifier is usedF)r   flattenr   )r	   	feat_mult)r   r   r   r   r   Zflatten_in_poolglobal_poolnum_pooled_features r   _/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/layers/classifier.py_create_pool   s    r   c                 C   s>   |dkrt  }n(|r*t j| |ddd}nt j| |dd}|S )Nr   r   T)bias)nnIdentityConv2dLinear)r   r   r   fcr   r   r   
_create_fc'   s    
r   NCHW)r   r   r   r   r   	drop_ratec           
      C   sH   t | ||||d\}}t|||d}|d ur@t|}	||	|fS ||fS )Nr   r   r   )r   r   r   Dropout)
r   r   r   r   r   r!   r   r   r   Zdropoutr   r   r   create_classifier1   s     


r%   c                       sL   e Zd ZdZdeeeeeed fddZdd
dZ	dedddZ
  ZS )ClassifierHeadz;Classifier head w/ configurable global pooling and dropout.r           Fr    )in_featuresr   r   r!   r   r   c           	         sn   t t|   || _|| _|| _t|||||d\}}|| _t	|| _
|| _|r`|r`tdnt | _dS )a.  
        Args:
            in_features: The number of input features.
            num_classes:  The number of classes for the final classifier layer (output).
            pool_type: Global pooling type, pooling disabled if empty string ('').
            drop_rate: Pre-classifier dropout rate.
        r"   r   N)superr&   __init__r(   r   r   r%   r   r   r$   dropr   Flattenr   r   )	selfr(   r   r   r!   r   r   r   r   	__class__r   r   r*   N   s    
zClassifierHead.__init__Nc                 C   sz   |d urT|| j jkrTt| j||| j| jd\| _ | _| jrH|rHtdnt	 | _
n"| j| j   }t||| jd| _d S )N)r   r   r   r   r#   )r   r   r%   r(   r   r   r   r   r,   r   r   r   r   )r-   r   r   r   r   r   r   reseto   s     zClassifierHead.reset
pre_logitsc                 C   s6   |  |}| |}|r"| |S | |}| |S N)r   r+   r   r   r-   xr2   r   r   r   forward   s    



zClassifierHead.forward)r   r'   Fr    )N)F)__name__
__module____qualname____doc__intstrfloatboolr*   r0   r6   __classcell__r   r   r.   r   r&   K   s       !
r&   c                
       s^   e Zd Zdeeee eeeeef eeef d fddZ	dd	d
Z
dedddZ  ZS )NormMlpClassifierHeadNr   r'   layernorm2dtanh)r(   r   hidden_sizer   r!   
norm_layer	act_layerc           	         s   t    || _|| _|| _| | _t|}t|}| jrHtt	j
ddnt	j}t|d| _||| _|rrt	dnt	 | _|rt	td|||fd| fg| _|| _n
t	 | _t	|| _|dkr|| j|nt	 | _dS )a  
        Args:
            in_features: The number of input features.
            num_classes:  The number of classes for the final classifier layer (output).
            hidden_size: The hidden size of the MLP (pre-logits FC layer) if not None.
            pool_type: Global pooling type, pooling disabled if empty string ('').
            drop_rate: Pre-classifier dropout rate.
            norm_layer: Normalization layer type.
            act_layer: MLP activation layer type (only used if hidden_size is not None).
        r   Zkernel_sizer   r   Zactr   N)r)   r*   r(   rC   r   r   r   r
   r   r   r   r   r	   r   normr,   r   r   Z
Sequentialr   r2   r$   r+   r   )	r-   r(   r   rC   r   r!   rD   rE   linear_layerr.   r   r   r*      s(    



zNormMlpClassifierHead.__init__c                 C   s  |d ur,t |d| _|r"tdnt | _| j | _| jrLttj	ddntj
}| jrt| jjtj	rn| jrt| jjtj
r| jrt T || j| j}|j| jjj|jj |j| jjj || j_W d    n1 s0    Y  |dkr|| j|nt | _d S )NrG   r   rF   r   )r	   r   r   r,   r   r   Zis_identityr   r   r   r   rC   
isinstancer2   r   torchZno_gradr(   ZweightZcopy_Zreshapeshaper   r   )r-   r   r   rI   Znew_fcr   r   r   r0      s"    
&zNormMlpClassifierHead.resetFr1   c                 C   sH   |  |}| |}| |}| |}| |}|r:|S | |}|S r3   )r   rH   r   r2   r+   r   r4   r   r   r   r6      s    





zNormMlpClassifierHead.forward)Nr   r'   rA   rB   )N)F)r7   r8   r9   r;   r   r<   r=   r   r   r*   r0   r>   r6   r?   r   r   r.   r   r@      s         

+
r@   )r   FN)F)r   Fr    N)r:   collectionsr   	functoolsr   typingr   r   r   rK   Ztorch.nnr   r   FZadaptive_avgmax_poolr	   Z
create_actr
   Zcreate_normr   r;   r<   r>   r   r   r=   r%   Moduler&   r@   r   r   r   r   <module>   sD      
    ?