a
    d3                     @   sx  d 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
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d*ddZG dd dejZG dd	 d	ejZdd Zd+ddZeeddedddedddedddedddZed,edd d!Zed-edd"d#Zed.edd$d%Z ed/edd&d'Z!ed0edd(d)Z"dS )1a  PyTorch SelecSLS Net example for ImageNet Classification
License: CC BY 4.0 (https://creativecommons.org/licenses/by/4.0/legalcode)
Author: Dushyant Mehta (@mehtadushy)

SelecSLS (core) Network Architecture as proposed in "XNect: Real-time Multi-person 3D
Human Pose Estimation with a Single RGB Camera, Mehta et al."
https://arxiv.org/abs/1907.00837

Based on ResNet implementation in https://github.com/rwightman/pytorch-image-models
and SelecSLS Net implementation in https://github.com/mehtadushy/SelecSLS-Pytorch
    )ListNIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)create_classifier   )build_model_with_cfg)register_modelgenerate_default_cfgsSelecSlsc                       sP   e Zd Z fddZejjdd Zejjdd Zeej	 dddZ  Z
S )SequentialListc                    s   t t| j|  d S N)superr   __init__)selfargs	__class__ ]/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/selecsls.pyr      s    zSequentialList.__init__c                 C   s   d S r   r   r   xr   r   r   forward   s    zSequentialList.forwardc                 C   s   d S r   r   r   r   r   r   r   $   s    returnc                 C   s   | D ]}||}q|S r   r   )r   r   moduler   r   r   r   )   s    
)__name__
__module____qualname__r   torchjit_overload_methodr   r   Tensor__classcell__r   r   r   r   r      s   

r   c                       sN   e Zd Zd
 fdd	Zejjdd Zejjdd Zejdd	dZ  Z	S )	SelectSeqindexr   c                    s   t t|   || _|| _d S r   )r   r$   r   moder%   )r   r&   r%   r   r   r   r   0   s    zSelectSeq.__init__c                 C   s   d S r   r   r   r   r   r   r   5   s    zSelectSeq.forwardc                 C   s   d S r   r   r   r   r   r   r   :   s    r   c                 C   s&   | j dkr|| j S tj|ddS d S )Nr%   r   )Zdim)r&   r%   r   catr   r   r   r   r   ?   s    

)r%   r   )
r   r   r   r   r   r    r!   r   r"   r#   r   r   r   r   r$   /   s   

r$      c                 C   sP   |d u r |d ||d   d }t t j| |||||ddt |t jddS )Nr      F)paddingdilationZbiasT)Zinplace)nn
SequentialConv2dZBatchNorm2dZReLU)in_chsout_chskstrider*   r+   r   r   r   conv_bnF   s    
r3   c                       s:   e Zd Zd fdd	Zeej eej dddZ  ZS )SelecSlsBlockr   c                    s   t t|   || _|| _|dv s&J t||d||d| _t||d| _t||d d| _t|d |d| _	t||d d| _
td| |rdn| |d| _d S )Nr   r)   r(   )r+   r   r)   r   )r   r4   r   r2   is_firstr3   conv1conv2conv3conv4conv5conv6)r   r/   Zskip_chsZmid_chsr0   r6   r2   r+   r   r   r   r   Q   s    zSelecSlsBlock.__init__)r   r   c              	   C   s   t |ts|g}t|dv s J | |d }| | |}| | |}| jrt| 	t
|||gd}||gS | 	t
||||d gd|d gS d S )Nr5   r   r   )
isinstancelistlenr7   r9   r8   r;   r:   r6   r<   r   r'   )r   r   d1Zd2Zd3outr   r   r   r   _   s    
zSelecSlsBlock.forward)r   )	r   r   r   r   r   r   r"   r   r#   r   r   r   r   r4   P   s   r4   c                       s   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edddZdd Z  ZS )r   a  SelecSls42 / SelecSls60 / SelecSls84

    Parameters
    ----------
    cfg : network config dictionary specifying block type, feature, and head args
    num_classes : int, default 1000
        Number of classification classes.
    in_chans : int, default 3
        Number of input (color) channels.
    drop_rate : float, default 0.
        Dropout probability before classifier, for training
    global_pool : str, default 'avg'
        Global pooling type. One of 'avg', 'max', 'avgmax', 'catavgmax'
      r(           avgc                    s   || _ tt|   t|ddd| _t fdd d D  | _t | _	t
jdd  d D  | _ d	 | _ d
 | _t| j| j ||d\| _| _| _|  D ](\}}t|t
jrt
jj|jddd qd S )N    r)   )r2   c                    s   g | ]} d  | qS )blockr   ).0Z
block_argscfgr   r   
<listcomp>       z%SelecSls.__init__.<locals>.<listcomp>featuresc                 S   s   g | ]}t | qS r   )r3   )rG   Z	conv_argsr   r   r   rJ      rK   headnum_featuresfeature_info)	pool_type	drop_rateZfan_outZrelu)r&   Znonlinearity)num_classesr   r   r   r3   stemr   rL   r$   from_seqr,   r-   rM   rN   rO   r   global_pool	head_dropfcZnamed_modulesr=   r.   initZkaiming_normal_Zweight)r   rI   rR   Zin_chansrQ   rU   nmr   rH   r   r   ~   s"    

zSelecSls.__init__Fc                 C   s   t ddddS )Nz^stemz^features\.(\d+)z^head)rS   blocksZblocks_head)dict)r   Zcoarser   r   r   group_matcher   s
    zSelecSls.group_matcherTc                 C   s   |rJ dd S )Nz$gradient checkpointing not supportedr   )r   enabler   r   r   set_grad_checkpointing   s    zSelecSls.set_grad_checkpointingc                 C   s   | j S r   )rW   )r   r   r   r   get_classifier   s    zSelecSls.get_classifierc                 C   s$   || _ t| j| j |d\| _| _d S )N)rP   )rR   r   rN   rU   rW   )r   rR   rU   r   r   r   reset_classifier   s    zSelecSls.reset_classifierc                 C   s(   |  |}| |}| | |}|S r   )rS   rL   rM   rT   r   r   r   r   forward_features   s    

zSelecSls.forward_features)
pre_logitsc                 C   s&   |  |}| |}|r|S | |S r   )rU   rV   rW   )r   r   rc   r   r   r   forward_head   s    

zSelecSls.forward_headc                 C   s   |  |}| |}|S r   )rb   rd   r   r   r   r   r      s    

zSelecSls.forward)rB   r(   rC   rD   )F)T)rD   )F)r   r   r   __doc__r   r   r    ignorer]   r_   r`   ra   rb   boolrd   r   r#   r   r   r   r   r   n   s   

c              	   K   sP  i }t ddddg}| drt|d< g d|d< |t d	d
ddt ddddt ddddg |t dddd | dkrg d|d< |t dddd d|d< n(g d|d< |t dddd d|d< nT| drt|d< g d|d< |t d	d
ddt ddddt ddd dg |t dddd | d!krhg d"|d< |t dddd d|d< n(g d#|d< |t dddd d|d< n| d$krt|d< g d%|d< |t d&d
ddt d'dd(dt d)dd*dg g d+|d< d|d< |t ddddt ddddg ntd,|  d- ||d.< tt| |f|t d/d0d1d2|S )3NrE   r)   zstem.2)Znum_chsZ	reductionr   
SelecSls42rF   )rE   r   @   rj   Tr)   rj   rj   rj      Fr   )rl   r      rm   Tr)   )rm   rm   rm      Fr   )rn   r   0  ro   Tr)   )ro   ro   ro     Fr   rL   rl      z
features.1rn      z
features.3rp      z
features.5   zhead.1SelecSls42b)rp     r(   r)   rw   rt   r(   r   rt      r(   r)   rz   rt   r   r   rM   rj   zhead.3rN   )rv   rx   rt   rt   r(   r)   rt   rz   r   r   rz   
SelecSls60)	ri   rk   )rl   r   rl   rl   Tr)   )rl   rl   rl   rl   Fr   )rl   rl   rl   rn   Fr   )rn   r   rn   rn   Tr)   rn   rn   rn   rn   Fr   r   )rn   rn   rn     Fr   z
features.4r   z
features.8SelecSls60b)r     r(   r)   r   rt   r(   r   ry   r{   )r   r   r|   r}   
SelecSls84)ri   )rj   rj   rj   rm   Fr   )rm   r   rm   rm   Tr)   rm   rm   rm   rm   Fr   r   r   )rm   rm   rm   ro   Fr   )ro   r   ro   ro   Tr)   ro   ro   ro   ro   Fr   r   r   r   )ro   ro   ro      Fr   rm   ro   z
features.6r   zfeatures.12))r   rw   r(   r)   rx   r|   )rt   rz   r(   r   zInvalid net configuration z !!!rO   )r   r   r)   r(   rq   T)Zout_indicesZflatten_sequential)Z	model_cfgZfeature_cfg)r\   
startswithr4   extendappend
ValueErrorr   r   )variant
pretrainedkwargsrI   rO   r   r   r   _create_SelecSls   sx    
	





r    c                 K   s   | dddddt tddd
|S )	NrB   )r(      r   )rq   rq   g      ?Zbilinearzstem.0rW   )
urlrR   Z
input_sizeZ	pool_sizeZcrop_pctinterpolationmeanZstdZ
first_conv
classifierr   )r   r   r   r   r   _cfg>  s    r   Zbicubic)r   ztimm/)Z	hf_hub_idr   )zSelecSls42.untrainedzSelecSls42b.in1kzSelecSls60.in1kzSelecSls60b.in1kzSelecSls84.untrainedFr   c                 K   s   t d| fi |S )z#Constructs a SelecSls42 model.
    rh   r   r   r   r   r   r   rh   Z  s    rh   c                 K   s   t d| fi |S )z%Constructs a SelecSls42_B model.
    ru   r   r   r   r   r   ru   a  s    ru   c                 K   s   t d| fi |S )z#Constructs a SelecSls60 model.
    r~   r   r   r   r   r   r~   h  s    r~   c                 K   s   t d| fi |S )z%Constructs a SelecSls60_B model.
    r   r   r   r   r   r   r   o  s    r   c                 K   s   t d| fi |S )z#Constructs a SelecSls84 model.
    r   r   r   r   r   r   r   v  s    r   )r(   r   Nr   )r   )F)F)F)F)F)#re   typingr   r   Ztorch.nnr,   Ztorch.nn.functionalZ
functionalFZ	timm.datar   r   Ztimm.layersr   Z_builderr   	_registryr	   r
   __all__r-   r   Moduler$   r3   r4   r   r   r   Zdefault_cfgsrh   ru   r~   r   r   r   r   r   r   <module>   s^   

K 
