a
    d=                  
   @   s  d Z ddlZddlmZ ddlZddlmZ ddlm  mZ	 ddl
m  mZ ddlmZ ddlmZmZ ddlmZmZmZmZ ddlmZ dd	lmZ dd
lmZmZ dg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d Z'dd Z(d.ddZ)ee)dddde)dddde) e)dde)dde)dde)dddZ*ed/e&d!d"d#Z+ed0e&d!d$d%Z,ed1e&d!d&d'Z-ed2e&d!d(d)Z.ed3e&d!d*d+Z/ed4e&d!d,d-Z0dS )5zPytorch Densenet implementation w/ tweaks
This file is a copy of https://github.com/pytorch/vision 'densenet.py' (BSD-3-Clause) with
fixed kwargs passthrough and addition of dynamic global avg/max pool.
    N)OrderedDict)ListIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)BatchNormAct2dget_norm_act_layer
BlurPool2dcreate_classifier   )build_model_with_cfg)MATCH_PREV_GROUP)register_modelgenerate_default_cfgsDenseNetc                       sl   e Zd Zeddf fdd	Zdd Zdd Zejj	d	d
 Z
ejjdd Zejjdd Zdd Z  ZS )
DenseLayer        Fc                    s   t t|   | d||f | dtj||| ddddf | d||| f | dtj|| |ddddd	f t|| _|| _d S )
Nnorm1conv1r   Fkernel_sizestridebiasnorm2conv2   r   r   paddingr   )	superr   __init__
add_modulennConv2dfloat	drop_rategrad_checkpointing)selfnum_input_featuresgrowth_ratebn_size
norm_layerr$   r%   	__class__ ]/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/densenet.pyr      s    	




zDenseLayer.__init__c                 C   s    t |d}| | |}|S Nr   )torchcatr   r   )r&   xsZconcated_featuresbottleneck_outputr-   r-   r.   bottleneck_fn+   s    zDenseLayer.bottleneck_fnc                 C   s   |D ]}|j r dS qdS )NTF)Zrequires_grad)r&   xZtensorr-   r-   r.   any_requires_grad2   s    zDenseLayer.any_requires_gradc                    s    fdd}t j|g|R  S )Nc                     s
     | S N)r4   )r2   r&   r-   r.   closure<   s    z6DenseLayer.call_checkpoint_bottleneck.<locals>.closure)cp
checkpoint)r&   r5   r9   r-   r8   r.   call_checkpoint_bottleneck9   s    z%DenseLayer.call_checkpoint_bottleneckc                 C   s   d S r7   r-   r&   r5   r-   r-   r.   forwardA   s    zDenseLayer.forwardc                 C   s   d S r7   r-   r=   r-   r-   r.   r>   F   s    c                 C   s   t |tjr|g}n|}| jrF| |rFtj r:td| |}n
| 	|}| 
| |}| jdkr~tj|| j| jd}|S )Nz%Memory Efficient not supported in JITr   )ptraining)
isinstancer0   ZTensorr%   r6   jitZis_scripting	Exceptionr<   r4   r   r   r$   FZdropoutr@   )r&   r5   Zprev_featuresr3   new_featuresr-   r-   r.   r>   M   s    


)__name__
__module____qualname__r   r   r4   r6   r0   rB   Zunusedr<   Z_overload_methodr>   __classcell__r-   r-   r+   r.   r      s   


r   c                       s0   e Zd ZdZeddf fdd	Zdd Z  ZS )
DenseBlock   r   Fc           
   	      sP   t t|   t|D ]4}t|||  |||||d}	| d|d  |	 qd S )N)r(   r)   r*   r$   r%   zdenselayer%dr   )r   rJ   r   ranger   r    )
r&   
num_layersr'   r)   r(   r*   r$   r%   ilayerr+   r-   r.   r   c   s    

zDenseBlock.__init__c                 C   s6   |g}|   D ]\}}||}|| qt|dS r/   )itemsappendr0   r1   )r&   Zinit_featuresfeaturesnamerO   rE   r-   r-   r.   r>   y   s
    zDenseBlock.forward)rF   rG   rH   _versionr   r   r>   rI   r-   r-   r+   r.   rJ   `   s   rJ   c                       s"   e Zd Zedf fdd	Z  ZS )DenseTransitionNc              
      sr   t t|   | d|| | dtj||dddd |d urX| d||dd n| dtjddd	 d S )
NZnormconvr   Fr   poolrK   )r   )r   r   )r   rU   r   r    r!   r"   Z	AvgPool2d)r&   r'   num_output_featuresr*   aa_layerr+   r-   r.   r      s    

zDenseTransition.__init__)rF   rG   rH   r   r   rI   r-   r-   r+   r.   rU      s   rU   c                       sp   e Zd 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d Z  ZS )!r   a/  Densenet-BC model class, based on
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`_

    Args:
        growth_rate (int) - how many filters to add each layer (`k` in paper)
        block_config (list of 4 ints) - how many layers in each pooling block
        bn_size (int) - multiplicative factor for number of bottle neck layers
          (i.e. bn_size * k features in the bottleneck layer)
        drop_rate (float) - dropout rate before classifier layer
        proj_drop_rate (float) - dropout rate after each dense layer
        num_classes (int) - number of classification classes
        memory_efficient (bool) - If True, uses checkpointing. Much more memory efficient,
          but slower. Default: *False*. See `"paper" <https://arxiv.org/pdf/1707.06990.pdf>`_
                      r   avg    relubatchnorm2dNr   FTc                    s2  || _ tt|   t|	|d}	d|v }|d }|
d u rJtjdddd}n"tjtjdddd|
|ddg }|r$| }}d|v rd|d	  }d
|v r|n
d|d	  }ttdtj	||dddddfd|	|fdtj	||dddddfd|	|fdtj	||dddddfd|	|fd|fg| _
n8ttdtj	||dddddfd|	|fd|fg| _
t|dd|rndnd dg| _d	}|}t|D ]\}}t|||||	||d}d|d  }| j
|| |||  }|rd n|
}|t|d kr|  jt||d| dg7  _|d9 }t||d |	|d}| j
d|d  | |d }q| j
d|	| |  jt||d dg7  _|| _t| j| j |d!\}}|| _t|| _|| _|  D ]r}t|tj	rtj|j nNt|tjrtj|jd tj|jd nt|tj rtj|jd qd S )"N)	act_layerdeeprK   r   r   )r   r   r   )Zchannelsr   Ztieredrb   Znarrowr\   Zconv0F)r   r   r   Znorm0r   r   r   r   Zpool0   r   zfeatures.normr   )Znum_chsZ	reductionmodule)rM   r'   r)   r(   r*   r$   r%   Z
denseblockz	features.)r'   rX   r*   rY   Z
transitionZnorm5zfeatures.norm5Z	pool_type)!num_classesr   r   r   r   r!   Z	MaxPool2d
Sequentialr   r"   rR   dictZfeature_info	enumeraterJ   r    lenrU   num_featuresr
   global_poolZDropout	head_drop
classifiermodulesrA   initZkaiming_normal_ZweightZBatchNorm2dZ	constant_r   ZLinear)r&   r(   block_configrk   Zin_chansrq   r)   	stem_typerf   r*   rY   r$   Zproj_drop_rateZmemory_efficientZaa_stem_onlyZ	deep_stemZnum_init_featuresZ	stem_poolZ
stem_chs_1Z
stem_chs_2Zcurrent_striderp   rN   rM   blockmodule_nameZtransition_aa_layerZtransrs   mr+   r-   r.   r      s    






	
zDenseNet.__init__c                 C   s    t d|rdn
ddtfgd}|S )Nz<^features\.conv[012]|features\.norm[012]|features\.pool[012]z)^features\.(?:denseblock|transition)(\d+))z+^features\.denseblock(\d+)\.denselayer(\d+)Nz^features\.transition(\d+))stemblocks)rm   r   )r&   ZcoarseZmatcherr-   r-   r.   group_matcher  s    zDenseNet.group_matcherc                 C   s$   | j  D ]}t|tr
||_q
d S r7   )rR   rt   rA   r   r%   )r&   enablebr-   r-   r.   set_grad_checkpointing  s    
zDenseNet.set_grad_checkpointingc                 C   s   | j S r7   )rs   r8   r-   r-   r.   get_classifier!  s    zDenseNet.get_classifierc                 C   s$   || _ t| j| j |d\| _| _d S )Nrj   )rk   r
   rp   rq   rs   )r&   rk   rq   r-   r-   r.   reset_classifier%  s    
zDenseNet.reset_classifierc                 C   s
   |  |S r7   )rR   r=   r-   r-   r.   forward_features*  s    zDenseNet.forward_featuresc                 C   s,   |  |}| |}| |}| |}|S r7   )r   rq   rr   rs   r=   r-   r-   r.   r>   -  s
    



zDenseNet.forward)rZ   r[   r`   r   ra   rb   rc   rd   re   Nr   r   FT)F)T)ra   )rF   rG   rH   __doc__r   r0   rB   ignorer}   r   r   r   r   r>   rI   r-   r-   r+   r.   r      s2                 m


c                 C   sT   t d}t|  D ]8}||}|r|d|d }| | | |< | |= q| S )Nz]^(.*denselayer\d+\.(?:norm|relu|conv))\.((?:[12])\.(?:weight|bias|running_mean|running_var))$r   rK   )recompilelistkeysmatchgroup)Z
state_dictpatternkeyresZnew_keyr-   r-   r.   _filter_torchvision_pretrained5  s    
r   c                 K   s0   ||d< ||d< t t| |ftddtd|S )Nr(   rv   T)Zflatten_sequential)Zfeature_cfgZpretrained_filter_fn)r   r   rm   r   )variantr(   rv   
pretrainedkwargsr-   r-   r.   _create_densenetB  s    r   rc   c                 K   s   | dddddt tddd
|S )	Nr`   )r      r   )rh   rh   g      ?Zbicubiczfeatures.conv0rs   )
urlrk   Z
input_sizeZ	pool_sizeZcrop_pctinterpolationmeanZstdZ
first_convrs   r   )r   r   r-   r-   r.   _cfgO  s    r   ztimm/)r      r   gffffff?)	hf_hub_idZtest_input_sizeZtest_crop_pct)r   )zdensenet121.ra_in1kzdensenetblur121d.ra_in1kzdensenet264d.untrainedzdensenet121.tv_in1kzdensenet169.tv_in1kzdensenet201.tv_in1kzdensenet161.tv_in1kF)returnc                 K   s   t ddd| d|}|S )ztDensenet-121 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    densenet121rZ   r[   r(   rv   r   )r   r   r   r   modelr-   r-   r.   r   g  s     r   c                 K   s   t ddd| dtd|}|S )zDensenet-121 w/ blur-pooling & 3-layer 3x3 stem
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    densenetblur121drZ   r[   rg   )r(   rv   r   rw   rY   )r   )r   r	   r   r-   r-   r.   r   q  s     r   c                 K   s   t ddd| d|}|S )ztDensenet-169 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    densenet169rZ   )r\   r]   rZ   rZ   r   )r   r   r   r-   r-   r.   r   |  s     r   c                 K   s   t ddd| d|}|S )ztDensenet-201 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    densenet201rZ   )r\   r]   0   rZ   r   )r   r   r   r-   r-   r.   r     s     r   c                 K   s   t ddd| d|}|S )ztDensenet-161 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    densenet161r   )r\   r]   $   r^   r   )r   r   r   r-   r-   r.   r     s     r   c                 K   s   t dddd| d|}|S )ztDensenet-264 model from
    `"Densely Connected Convolutional Networks" <https://arxiv.org/pdf/1608.06993.pdf>`
    densenet264dr   )r\   r]   @   r   rg   )r(   rv   rw   r   )r   r   r   r-   r-   r.   r     s     r   )rc   )F)F)F)F)F)F)1r   r   collectionsr   r0   Ztorch.nnr!   Ztorch.nn.functionalZ
functionalrD   Ztorch.utils.checkpointutilsr;   r:   Ztorch.jit.annotationsr   Z	timm.datar   r   Ztimm.layersr   r   r	   r
   Z_builderr   Z_manipulater   	_registryr   r   __all__Moduler   Z
ModuleDictrJ   rl   rU   r   r   r   r   Zdefault_cfgsr   r   r   r   r   r   r-   r-   r-   r.   <module>   s`   I! #
		
			