a
    dv                      @   s^  d Z ddlmZ ddlmZ ddlZddlmZ ddlm	Z	m
Z
 ddlmZmZmZmZmZmZmZmZmZmZmZmZmZ ddlmZ dd	lmZmZmZ dd
lm Z 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(G dd dej$Z)dd Z*dddeeeddfddZ+G dd dej$Z,duej$e-d d!d"Z.e/ dvej$e-e-d$d%d&Z0dwd(d)Z1dxd*d+Z2dyd,d-Z3e e3d.d/dd0e3d.d/dd0e3d.d1d2d3d/dd4e3d.d5d6d3dd7e3d.d5d6d3dd7e3d.d5d6d3dd7e3d.d5d6d3dd7e3d.d5d6d3dd7e3d.d8d9d3dd7e3d.d:dd;e3d.d:dd;e3d.d:dd;e3d.d:dd;e3d.d:dd;e3d.d:dd;e3d.d/d<d=d3d>e3d/d?d@e3d/d?d@e3d.d/d<d=d3d>e3d/d?d@e3d/dAe3d/d?d@e3d.d/d?d<d=d3dBe3d.d/d?d<d=d3dBe3d/d?d@dCZ4e!dze,dDdEdFZ5e!d{e,dDdGdHZ6e!d|e,dDdIdJZ7e!d}e,dDdKdLZ8e!d~e,dDdMdNZ9e!de,dDdOdPZ:e!de,dDdQdRZ;e!de,dDdSdTZ<e!de,dDdUdVZ=e!de,dDdWdXZ>e!de,dDdYdZZ?e!de,dDd[d\Z@e!de,dDd]d^ZAe!de,dDd_d`ZBe!de,dDdadbZCe!de,dDdcddZDe"eEdedfdgdhdidjdkdldmdndodpdqdrdsdt dS )a/  Pre-Activation ResNet v2 with GroupNorm and Weight Standardization.

A PyTorch implementation of ResNetV2 adapted from the Google Big-Transfoer (BiT) source code
at https://github.com/google-research/big_transfer to match timm interfaces. The BiT weights have
been included here as pretrained models from their original .NPZ checkpoints.

Additionally, supports non pre-activation bottleneck for use as a backbone for Vision Transfomers (ViT) and
extra padding support to allow porting of official Hybrid ResNet pretrained weights from
https://github.com/google-research/vision_transformer

Thanks to the Google team for the above two repositories and associated papers:
* Big Transfer (BiT): General Visual Representation Learning - https://arxiv.org/abs/1912.11370
* An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale - https://arxiv.org/abs/2010.11929
* Knowledge distillation: A good teacher is patient and consistent - https://arxiv.org/abs/2106.05237

Original copyright of Google code below, modifications by Ross Wightman, Copyright 2020.
    )OrderedDict)partialNIMAGENET_INCEPTION_MEANIMAGENET_INCEPTION_STD)GroupNormActBatchNormAct2dEvoNorm2dS0FilterResponseNormTlu2dClassifierHeadDropPathAvgPool2dSamecreate_pool2d	StdConv2dcreate_conv2dget_act_layerget_norm_act_layermake_divisible   )build_model_with_cfg)checkpoint_seqnamed_applyadapt_input_conv)generate_default_cfgsregister_modelregister_model_deprecationsResNetV2c                       s2   e Zd ZdZd fdd	Zdd	 Zd
d Z  ZS )PreActBottlenecka  Pre-activation (v2) bottleneck block.

    Follows the implementation of "Identity Mappings in Deep Residual Networks":
    https://github.com/KaimingHe/resnet-1k-layers/blob/master/resnet-pre-act.lua

    Except it puts the stride on 3x3 conv when available.
    N      ?r           c              
      s   t    |p|}|	pt}	|
p(ttdd}
|p0|}t|| }|d urb||||||d|	|
d| _nd | _|
|| _|	||d| _|
|| _	|	||d|||d| _
|
|| _|	||d| _|dkrt|nt | _d S )	N    Z
num_groupsT)stridedilationfirst_dilationpreact
conv_layer
norm_layerr      r"   r#   groupsr   )super__init__r   r   r   r   
downsamplenorm1conv1norm2conv2norm3conv3r   nnIdentity	drop_pathselfin_chsout_chsbottle_ratior"   r#   r$   r*   	act_layerr&   r'   
proj_layerdrop_path_rateZmid_chs	__class__ ]/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/resnetv2.pyr,   :   s&    




zPreActBottleneck.__init__c                 C   s   t j| jj d S N)r4   initzeros_r3   weightr8   rA   rA   rB   zero_init_last_   s    zPreActBottleneck.zero_init_lastc                 C   s^   |  |}|}| jd ur"| |}| |}| | |}| | |}| |}|| S rC   )r.   r-   r/   r1   r0   r3   r2   r6   )r8   xZx_preactshortcutrA   rA   rB   forwardb   s    




zPreActBottleneck.forward)Nr   r   r   Nr   NNNNr   __name__
__module____qualname____doc__r,   rH   rK   __classcell__rA   rA   r?   rB   r   1   s              %r   c                       s2   e Zd ZdZd fdd	Zdd	 Zd
d Z  ZS )
BottleneckzUNon Pre-activation bottleneck block, equiv to V1.5/V1b Bottleneck. Used for ViT.
    Nr   r   r   c              	      s   t    |p|}|ptj}|	p"t}	|
p2ttdd}
|p:|}t|| }|d urj|||||d|	|
d| _nd | _|	||d| _	|
|| _
|	||d|||d| _|
|| _|	||d| _|
|dd| _|d	krt|nt | _|d
d| _d S )Nr    r!   F)r"   r#   r%   r&   r'   r   r(   r)   Z	apply_actr   T)Zinplace)r+   r,   r4   ReLUr   r   r   r   r-   r/   r.   r1   r0   r3   r2   r   r5   r6   act3r7   r?   rA   rB   r,   u   s*    





zBottleneck.__init__c                 C   s&   t | jdd d ur"tj| jj d S )NrF   )getattrr2   r4   rD   rE   rF   rG   rA   rA   rB   rH      s    zBottleneck.zero_init_lastc                 C   sp   |}| j d ur|  |}| |}| |}| |}| |}| |}| |}| |}| || }|S rC   )	r-   r/   r.   r1   r0   r3   r2   r6   rU   )r8   rI   rJ   rA   rA   rB   rK      s    








zBottleneck.forward)Nr   r   r   Nr   NNNNr   rL   rA   rA   r?   rB   rR   r   s              'rR   c                       s&   e Zd Zd fdd	Zdd Z  ZS )	DownsampleConvr   NTc	           	         s>   t t|   |||d|d| _|r,t n
||dd| _d S )Nr   r"   FrS   )r+   rW   r,   convr4   r5   norm)	r8   r9   r:   r"   r#   r$   r%   r&   r'   r?   rA   rB   r,      s    zDownsampleConv.__init__c                 C   s   |  | |S rC   )rZ   rY   r8   rI   rA   rA   rB   rK      s    zDownsampleConv.forward)r   r   NTNNrM   rN   rO   r,   rK   rQ   rA   rA   r?   rB   rW      s         rW   c                       s&   e Zd Zd fdd	Zdd Z  ZS )	DownsampleAvgr   NTc	                    s   t t|   |dkr|nd}	|dks.|dkr\|	dkrB|dkrBtntj}
|
d|	ddd| _n
t | _|||ddd| _|rt n
||dd| _	dS )	zd AvgPool Downsampling as in 'D' ResNet variants. This is not in RegNet space but I might experiment.r      TF)Z	ceil_modeZcount_include_padrX   rS   N)
r+   r]   r,   r   r4   Z	AvgPool2dpoolr5   rY   rZ   )r8   r9   r:   r"   r#   r$   r%   r&   r'   Z
avg_strideZavg_pool_fnr?   rA   rB   r,      s    
zDownsampleAvg.__init__c                 C   s   |  | | |S rC   )rZ   rY   r_   r[   rA   rA   rB   rK      s    zDownsampleAvg.forward)r   r   NTNNr\   rA   rA   r?   rB   r]      s         r]   c                       s:   e Zd ZdZddddedddf fdd	Zdd	 Z  ZS )
ResNetStagezResNet Stage.r   r   FNc                    s   t t|   |dv rdnd}t|||d}|r4tnt}|}t | _t	|D ]d}|	r^|	| nd}|dkrn|nd}| j
t||
||f|||||||d|| |}|}d }qNd S )N)r   r^   r   r^   )r<   r&   r'   r   r   )r"   r#   r;   r*   r$   r=   r>   )r+   r`   r,   dictr]   rW   r4   
Sequentialblocksrange
add_modulestr)r8   r9   r:   r"   r#   depthr;   r*   avg_down	block_dprblock_fnr<   r&   r'   Zblock_kwargsr$   Zlayer_kwargsr=   prev_chsZ	block_idxr>   r?   rA   rB   r,      s8    

zResNetStage.__init__c                 C   s   |  |}|S rC   )rc   r[   rA   rA   rB   rK     s    
zResNetStage.forward)rM   rN   rO   rP   r   r,   rK   rQ   rA   rA   r?   rB   r`      s   +r`   c                    s   t  fdddD S )Nc                    s   g | ]}| v qS rA   rA   ).0s	stem_typerA   rB   
<listcomp>      z is_stem_deep.<locals>.<listcomp>)deeptiered)anyrn   rA   rn   rB   is_stem_deep  s    ru   @    Tr    r!   c                 C   sX  t  }|dv sJ t|rd|v r8d| d |d f}n|d |d f}|| |d ddd|d< ||d |d	< ||d |d
 dd
d|d< ||d
 |d< ||d
 |dd
d|d< |s|||d< n$|| |ddd|d< |s|||d< d|v rtd
d|d< tjdddd|d< n4d|v r:tddddd|d< ntjddd
d|d< t|S )N)rw   fixedsamerr   Z
deep_fixedZ	deep_samers   rs   r(      r^   r   )kernel_sizer"   r/   r.   r   r1   r0   r3   r2      rY   rZ   rx   r   pad)r{   r"   paddingr_   ry   max)r   ru   r4   ZConstantPad2dZ	MaxPool2dr   rb   )r9   r:   ro   r%   r&   r'   stemstem_chsrA   rA   rB   create_resnetv2_stem  s.    

r   c                       s   e Zd ZdZdddddddd	d
dejeeddeddd
f fdd	Z	e
jjd$ddZe
j d%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   z7Implementation of Pre-activation (v2) ResNet mode.
    )   i   i   i     r(   avgr    r   rv   rw   FTr!   r   c                     s  t    || _|| _|}t||d}t|}g | _t|| }t|||	|||d| _	|rjt
|	rfdqldnd}| jt|d|d |}d}d	}d
d td|t||D }|rtnt}t | _tt|||D ]\}\}}}t|| }|dkrd	nd}||kr||9 }d	}t||||||
|||||d}|}||9 }|  jt||d| dg7  _| jt|| q|| _|r|| jnt | _t| j||| jdd| _| j |d d| _!dS )a  
        Args:
            layers (List[int]) : number of layers in each block
            channels (List[int]) : number of channels in each block:
            num_classes (int): number of classification classes (default 1000)
            in_chans (int): number of input (color) channels. (default 3)
            global_pool (str): Global pooling type. One of 'avg', 'max', 'avgmax', 'catavgmax' (default 'avg')
            output_stride (int): output stride of the network, 32, 16, or 8. (default 32)
            width_factor (int): channel (width) multiplication factor
            stem_chs (int): stem width (default: 64)
            stem_type (str): stem type (default: '' == 7x7)
            avg_down (bool): average pooling in residual downsampling (default: False)
            preact (bool): pre-activiation (default: True)
            act_layer (Union[str, nn.Module]): activation layer
            norm_layer (Union[str, nn.Module]): normalization layer
            conv_layer (nn.Module): convolution module
            drop_rate: classifier dropout rate (default: 0.)
            drop_path_rate: stochastic depth rate (default: 0.)
            zero_init_last: zero-init last weight in residual path (default: False)
        )r<   )r&   r'   z
stem.conv3	stem.convz	stem.normr^   )Znum_chsZ	reductionmodule   r   c                 S   s   g | ]}|  qS rA   )tolist)rl   rI   rA   rA   rB   rp     rq   z%ResNetV2.__init__.<locals>.<listcomp>r   )	r"   r#   rg   rh   r<   r&   r'   ri   rj   zstages.T)Z	pool_type	drop_rateZuse_convrH   FN)"r+   r,   num_classesr   r   r   Zfeature_infor   r   r   ru   appendra   torchZlinspacesumsplitr   rR   r4   rb   stages	enumeratezipr`   re   rf   Znum_featuresr5   rZ   r   headinit_weightsgrad_checkpointing) r8   layersZchannelsr   Zin_chansglobal_poolZoutput_stridewidth_factorr   ro   rh   r%   r<   r'   r&   r   r>   rH   ZwfZ	stem_featrk   Zcurr_strider#   Z
block_dprsrj   Z	stage_idxdcZbdprr:   r"   stager?   rA   rB   r,   H  st    (
"

 zResNetV2.__init__c                 C   s   t tt|d|  d S )Nr   )r   r   _init_weights)r8   rH   rA   rA   rB   r     s    zResNetV2.init_weightsresnet/c                 C   s   t | || d S rC   )_load_weights)r8   checkpoint_pathprefixrA   rA   rB   load_pretrained  s    zResNetV2.load_pretrainedc                 C   s   t d|rdnddgd}|S )Nz^stemz^stages\.(\d+))z^stages\.(\d+)\.blocks\.(\d+)N)z^norm)i )r   rc   )ra   )r8   ZcoarseZmatcherrA   rA   rB   group_matcher  s    zResNetV2.group_matcherc                 C   s
   || _ d S rC   )r   )r8   enablerA   rA   rB   set_grad_checkpointing  s    zResNetV2.set_grad_checkpointingc                 C   s   | j jS rC   )r   fcrG   rA   rA   rB   get_classifier  s    zResNetV2.get_classifierc                 C   s   || _ | j|| d S rC   )r   r   reset)r8   r   r   rA   rA   rB   reset_classifier  s    zResNetV2.reset_classifierc                 C   sD   |  |}| jr,tj s,t| j|dd}n
| |}| |}|S )NT)flatten)r   r   r   jitZis_scriptingr   r   rZ   r[   rA   rA   rB   forward_features  s    


zResNetV2.forward_features
pre_logitsc                 C   s   | j ||dS )Nr   )r   )r8   rI   r   rA   rA   rB   forward_head  s    zResNetV2.forward_headc                 C   s   |  |}| |}|S rC   )r   r   r[   rA   rA   rB   rK     s    

zResNetV2.forward)T)r   )F)T)r   )F)rM   rN   rO   rP   r4   rT   r   r   r   r,   r   r   ignorer   r   r   r   r   r   r   boolr   rK   rQ   rA   rA   r?   rB   r   D  s@   
g


	)r   namec                 C   s   t | tjs d|v rDt | tjrDtjj| jddd tj| j nt | tjr~tjj	| jddd | jd urtj| j nJt | tj
tjtjfrtj| j tj| j n|rt| dr|   d S )	Nhead.fcr   g{Gz?)meanstdZfan_outZrelu)modeZnonlinearityrH   )
isinstancer4   ZLinearConv2drD   Znormal_rF   rE   biasZkaiming_normal_ZBatchNorm2dZ	LayerNormZ	GroupNormZones_hasattrrH   )r   r   rH   rA   rA   rB   r     s     
r   r   )modelr   r   c              
   C   s  dd l }dd }||}t| jjjjd ||| d }| jjj| | jj||| d  | jj	||| d  t
t| jdd tjr| jjjjd || d	 jd
 kr| jjj||| d	  | jjj	||| d  t| j D ]\}\}}	t|	j D ]t\}
\}}d}| d|d  d|
d dd}|jj||| d| d  |jj||| d| d  |jj||| d| d  |jj||| d  |jj||| d  |jj||| d  |jj	||| d  |jj	||| d  |jj	||| d  |jd ur&|| d| d }|jjj|| q&q
d S )Nr   c                 S   s"   | j dkr| g d} t| S )zPossibly convert HWIO to OIHW.r   )r(   r^   r   r   )ndimZ	transposer   Z
from_numpy)Zconv_weightsrA   rA   rB   t2p  s    
z_load_weights.<locals>.t2pr   z%root_block/standardized_conv2d/kernelzgroup_norm/gammazgroup_norm/betar   zhead/conv2d/kernelzhead/conv2d/biasZstandardized_conv2dblockz/unitZ02d/za/z/kernelzb/zc/za/group_norm/gammazb/group_norm/gammazc/group_norm/gammaza/group_norm/betazb/group_norm/betazc/group_norm/betaza/proj/)numpyloadr   r   rY   rF   shapeZcopy_rZ   r   r   rV   r   r4   r   r   r   r   Znamed_childrenrc   r/   r1   r3   r.   r0   r2   r-   )r   r   r   npr   weightsZstem_conv_wiZsnamer   jZbnamer   cnameZblock_prefixwrA   rA   rB   r     s<    
" """r   Fc                 K   s"   t dd}tt| |fd|i|S )NT)Zflatten_sequentialfeature_cfg)ra   r   r   )variant
pretrainedkwargsr   rA   rA   rB   _create_resnetv2  s    
r   c                 K   s    t | f|dttddd|S )Nrx   g:0yE>)Zeps)r   ro   r&   )r   r   r   )r   r   r   rA   rA   rB   _create_resnetv2_bit  s    
r   c                 K   s   | dddddt tddd
|S )	Nr   )r(      r   )r|   r|   g      ?Zbilinearr   r   )
urlr   
input_size	pool_sizecrop_pctinterpolationr   r   
first_conv
classifierr   )r   r   rA   rA   rB   _cfg(  s    r   ztimm/Zbicubic)	hf_hub_idr   custom_load)r(     r   )   r   g      ?)r   r   r   r   r   r   )r(     r   )   r   )r   r   r   r   r   )r(     r   )   r   iSU  )r   r   r   gffffff?)r(      r   )r   r   r   test_input_sizetest_crop_pctz
stem.conv1)r   r   )r   )r   r   r   r   r   r   )%resnetv2_50x1_bit.goog_distilled_in1k-resnetv2_152x2_bit.goog_teacher_in21k_ft_in1k1resnetv2_152x2_bit.goog_teacher_in21k_ft_in1k_384$resnetv2_50x1_bit.goog_in21k_ft_in1k$resnetv2_50x3_bit.goog_in21k_ft_in1k%resnetv2_101x1_bit.goog_in21k_ft_in1k%resnetv2_101x3_bit.goog_in21k_ft_in1k%resnetv2_152x2_bit.goog_in21k_ft_in1k%resnetv2_152x4_bit.goog_in21k_ft_in1kresnetv2_50x1_bit.goog_in21kresnetv2_50x3_bit.goog_in21kresnetv2_101x1_bit.goog_in21kresnetv2_101x3_bit.goog_in21kresnetv2_152x2_bit.goog_in21kresnetv2_152x4_bit.goog_in21kzresnetv2_50.a1h_in1kzresnetv2_50d.untrainedzresnetv2_50t.untrainedzresnetv2_101.a1h_in1kzresnetv2_101d.untrainedzresnetv2_152.untrainedzresnetv2_152d.untrainedzresnetv2_50d_gn.ah_in1kzresnetv2_50d_evos.ah_in1kzresnetv2_50d_frn.untrained)returnc                 K   s   t d| g ddd|S )Nresnetv2_50x1_bitr(   r      r(   r   r   r   r   )r   r   r   r   rA   rA   rB   r     s     
r   c                 K   s   t d| g ddd|S )Nresnetv2_50x3_bitr   r(   r   )r   r   r   rA   rA   rB   r     s     
r   c                 K   s   t d| g ddd|S )Nresnetv2_101x1_bitr(   r      r(   r   r   )r   r   r   rA   rA   rB   r     s     
r   c                 K   s   t d| g ddd|S )Nresnetv2_101x3_bitr   r(   r   )r   r   r   rA   rA   rB   r     s     
r   c                 K   s   t d| g ddd|S )Nresnetv2_152x2_bitr(   rz   $   r(   r^   r   )r   r   r   rA   rA   rB   r     s     
r   c                 K   s   t d| g ddd|S )Nresnetv2_152x4_bitr   r   r   )r  r   r   rA   rA   rB   r    s     
r  c                 K   s0   t g dttd}tdd| it |fi |S )Nr   r   r&   r'   resnetv2_50r   )r  ra   r   r   r   r   r   Z
model_argsrA   rA   rB   r    s    r  c                 K   s4   t g dttddd}tdd| it |fi |S )Nr   rr   Tr   r&   r'   ro   rh   resnetv2_50dr   )r  r  r  rA   rA   rB   r    s
    
r  c                 K   s4   t g dttddd}tdd| it |fi |S )Nr   rs   Tr  resnetv2_50tr   )r	  r  r  rA   rA   rB   r	    s
    
r	  c                 K   s0   t g dttd}tdd| it |fi |S )Nr   r  resnetv2_101r   )r
  r  r  rA   rA   rB   r
    s    r
  c                 K   s4   t g dttddd}tdd| it |fi |S )Nr   rr   Tr  resnetv2_101dr   )r  r  r  rA   rA   rB   r    s
    
r  c                 K   s0   t g dttd}tdd| it |fi |S )Nr   r  resnetv2_152r   )r  r  r  rA   rA   rB   r    s    r  c                 K   s4   t g dttddd}tdd| it |fi |S )Nr   rr   Tr  resnetv2_152dr   )r  r  r  rA   rA   rB   r    s
    
r  c                 K   s4   t g dttddd}tdd| it |fi |S )Nr   rr   Tr  resnetv2_50d_gnr   )r  )ra   r   r   r   r  rA   rA   rB   r    s
    
r  c                 K   s4   t g dttddd}tdd| it |fi |S )Nr   rr   Tr  resnetv2_50d_evosr   )r  )ra   r   r	   r   r  rA   rA   rB   r    s
    
r  c                 K   s4   t g dttddd}tdd| it |fi |S )Nr   rr   Tr  resnetv2_50d_frnr   )r  )ra   r   r
   r   r  rA   rA   rB   r    s
    
r  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   )Zresnetv2_50x1_bitmZresnetv2_50x3_bitmZresnetv2_101x1_bitmZresnetv2_101x3_bitmZresnetv2_152x2_bitmZresnetv2_152x4_bitmZresnetv2_50x1_bitm_in21kZresnetv2_50x3_bitm_in21kZresnetv2_101x1_bitm_in21kZresnetv2_101x3_bitm_in21kZresnetv2_152x2_bitm_in21kZresnetv2_152x4_bitm_in21kZresnetv2_50x1_bit_distilledZresnetv2_152x2_bit_teacherZresnetv2_152x2_bit_teacher_384)rw   T)r   )F)F)rw   )F)F)F)F)F)F)F)F)F)F)F)F)F)F)F)F)FrP   collectionsr   	functoolsr   r   Ztorch.nnr4   Z	timm.datar   r   Ztimm.layersr   r   r	   r
   r   r   r   r   r   r   r   r   r   Z_builderr   Z_manipulater   r   r   	_registryr   r   r   __all__Moduler   rR   rW   r]   r`   ru   r   r   rf   r   Zno_gradr   r   r   r   Zdefault_cfgsr   r   r   r   r   r  r  r  r	  r
  r  r  r  r  r  r  rM   rA   rA   rA   rB   <module>   sb  <A@2

- %
	



R	