a
    dS                  	   @   sV  d Z ddlZddlZddlZddlmZ ddl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mZmZ ddlmZ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dl"m#Z#m$Z$m%Z% dgZ&e'e(Z)G dd de	j*Z+G dd de	j*Z,G dd de	j*Z-e.dddZ/ee.dddZ0G dd de	j*Z1G dd de	j*Z2d;e	j*e3e4ddd Z5d!d" Z6d#d$ Z7d<d&d'Z8d=d(d)Z9e$e9 e9 e9 e9d*d+e9d*d+e9d*d+d,Z:e#d>e2d-d.d/Z;e#d?e2d-d0d1Z<e#d@e2d-d2d3Z=e#dAe2d-d4d5Z>e#dBe2d-d6d7Z?e#dCe2d-d8d9Z@e%e(d5d7d9d: dS )Da   Nested Transformer (NesT) in PyTorch

A PyTorch implement of Aggregating Nested Transformers as described in:

'Aggregating Nested Transformers'
    - https://arxiv.org/abs/2105.12723

The official Jax code is released and available at https://github.com/google-research/nested-transformer. The weights
have been converted with convert/convert_nest_flax.py

Acknowledgments:
* The paper authors for sharing their research, code, and model weights
* Ross Wightman's existing code off which I based this

Copyright 2021 Alexander Soare
    N)partial)nnIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)
PatchEmbedMlpDropPathcreate_classifiertrunc_normal__assert)create_conv2dcreate_pool2d	to_ntupleuse_fused_attn	LayerNorm   )build_model_with_cfg)register_notrace_function)checkpoint_seqnamed_apply)register_modelgenerate_default_cfgsregister_model_deprecationsNestc                       s<   e Zd ZU dZejje ed< d
 fdd	Z	dd	 Z
  ZS )	Attentionz
    This is much like `.vision_transformer.Attention` but uses *localised* self attention by accepting an input with
     an extra "image block" dim
    
fused_attn   F        c                    sj   t    || _|| }|d | _t | _tj|d| |d| _t	|| _
t||| _t	|| _d S )Ng         )bias)super__init__	num_headsscaler   r   r   LinearqkvDropout	attn_dropproj	proj_drop)selfdimr#   qkv_biasr(   r*   Zhead_dim	__class__ Y/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/nest.pyr"   /   s    

zAttention.__init__c              	   C   s   |j \}}}}| ||||d| j|| j dddddd}|d\}}}	| jrntj|||	| j	j
d}n8|| j }||dd	 }
|
jd	d
}
| 	|
}
|
|	 }|ddddd||||}| |}| |}|S )zm
        x is shape: B (batch_size), T (image blocks), N (seq length per image block), C (embed dim)
        r   r      r         )Z	dropout_p)r,   )shaper&   reshaper#   permuteZunbindr   FZscaled_dot_product_attentionr(   pr$   	transposeZsoftmaxr)   r*   )r+   xBTNCr&   qkvattnr0   r0   r1   forward;   s    2



zAttention.forward)r   Fr   r   )__name__
__module____qualname____doc__torchjitFinalbool__annotations__r"   rF   __classcell__r0   r0   r.   r1   r   (   s   
r   c                       s<   e Zd ZdZdddddejejf fdd	Zdd Z  Z	S )	TransformerLayerz
    This is much like `.vision_transformer.Block` but:
        - Called TransformerLayer here to allow for "block" as defined in the paper ("non-overlapping image blocks")
        - Uses modified Attention layer that handles the "block" dimension
          @Fr   c
                    sn   t    |	|| _t|||||d| _|dkr8t|nt | _|	|| _	t
|| }
t||
||d| _d S )N)r#   r-   r(   r*   r   )Zin_featuresZhidden_features	act_layerZdrop)r!   r"   norm1r   rE   r	   r   Identity	drop_pathnorm2intr   mlp)r+   r,   r#   	mlp_ratior-   r*   r(   rV   rS   
norm_layerZmlp_hidden_dimr.   r0   r1   r"   Z   s$    


zTransformerLayer.__init__c                 C   s<   |  |}|| | | }|| | | | }|S N)rT   rV   rE   rY   rW   )r+   r=   yr0   r0   r1   rF   y   s    
zTransformerLayer.forward)
rG   rH   rI   rJ   r   GELUr   r"   rF   rP   r0   r0   r.   r1   rQ   T   s   	rQ   c                       s&   e Zd Zd fdd	Zdd Z  ZS )ConvPool c                    s>   t    t||d|dd| _||| _tddd|d| _d S )Nr   T)kernel_sizepaddingr    maxr3   )ra   Zstriderb   )r!   r"   r   convnormr   pool)r+   Zin_channelsZout_channelsr[   pad_typer.   r0   r1   r"      s    

zConvPool.__init__c                 C   sj   t |jd d dkd t |jd d dkd | |}| |dddddddd}| |}|S )z:
        x is expected to have shape (B, C, H, W)
        r5   r3   r   z1BlockAggregation requires even input spatial dimsr6   r   r   )r   r7   rd   re   r9   rf   r+   r=   r0   r0   r1   rF      s    
"
zConvPool.forward)r`   )rG   rH   rI   r"   rF   rP   r0   r0   r.   r1   r_      s   r_   )
block_sizec                 C   sv   | j \}}}}t|| dkd t|| dkd || }|| }| ||||||} | dd||| d|} | S )zimage to blocks
    Args:
        x (Tensor): with shape (B, H, W, C)
        block_size (int): edge length of a single square block in units of H, W
    r   z,`block_size` must divide input height evenlyz+`block_size` must divide input width evenlyr3   r   r6   )r7   r   r8   r<   )r=   ri   r>   HWrA   Zgrid_heightZ
grid_widthr0   r0   r1   blockify   s    rl   c           	      C   sX   | j \}}}}tt|}||  }}| ||||||} | dd||||} | S )zblocks to image
    Args:
        x (Tensor): with shape (B, T, N, C) where T is number of blocks and N is sequence size per block
        block_size (int): edge length of a single square block in units of desired H, W
    r3   r   )r7   rX   mathsqrtr8   r<   )	r=   ri   r>   r?   _rA   Z	grid_sizeheightwidthr0   r0   r1   
deblockify   s    rr   c                	       s<   e Zd ZdZdddddg dddf	 fdd	Zd	d
 Z  ZS )	NestLevelz7 Single hierarchical level of a Nested Transformer
    NrR   Tr   r`   c              
      s   t    || _d| _ttd||| _|d urJt	||d| _
n
t | _
trpt|kspJ dtj f	ddt|D  | _d S )NFr   )r[   rg   zDMust provide as many drop path rates as there are transformer layersc                    s*   g | ]"}t |  d 	qS ))	r,   r#   rZ   r-   r*   r(   rV   r[   rS   )rQ   .0i	rS   r(   rV   	embed_dimrZ   r[   r#   r*   r-   r0   r1   
<listcomp>   s   z&NestLevel.__init__.<locals>.<listcomp>)r!   r"   ri   grad_checkpointingr   	ParameterrK   zeros	pos_embedr_   rf   rU   len
Sequentialrangetransformer_encoder)r+   
num_blocksri   
seq_lengthr#   depthrx   Zprev_embed_dimrZ   r-   r*   r(   rV   r[   rS   rg   r.   rw   r1   r"      s    

zNestLevel.__init__c                 C   st   |  |}|dddd}t|| j}|| j }| jrNtj sNt	| j
|}n
| 
|}t|| j}|ddddS )z+
        expects x as (B, C, H, W)
        r   r3   r   r   )rf   r9   rl   ri   r}   rz   rK   rL   Zis_scriptingr   r   rr   rh   r0   r0   r1   rF      s    


zNestLevel.forward)rG   rH   rI   rJ   r"   rF   rP   r0   r0   r.   r1   rs      s   
.rs   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 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   z Nested Transformer (NesT)

    A PyTorch impl of : `Aggregating Nested Transformers`
        - https://arxiv.org/abs/2105.12723
       r   r2         i   r2   r      r3   r3        rR   Tr         ?Nr`   avgc                    s  t    dD ]8}t | }t|tjjrt||ksJ d| dqt||}t||}t||}|| _	|d | _
g | _|pt}|ptj}|| _|| _t|tjjr|d |d ksJ d|d }|| dksJ d|| _d	t| d | _|| t| jd  dks(J d
t|| t| jd  | _t||||d dd| _| jj| _| j| jd  | _g }dd td|t |!|D }d}d	}t"t| jD ]|}|| }|#t$| j| | j| j|| || |||	|
|||| |||d |  jt%||d| dg7  _|}|d9 }qtj&| | _'||d | _(t)| j
| j	|d\}}|| _*t+|| _,|| _-| .| dS )a  
        Args:
            img_size (int, tuple): input image size
            in_chans (int): number of input channels
            patch_size (int): patch size
            num_levels (int): number of block hierarchies (T_d in the paper)
            embed_dims (int, tuple): embedding dimensions of each level
            num_heads (int, tuple): number of attention heads for each level
            depths (int, tuple): number of transformer layers for each level
            num_classes (int): number of classes for classification head
            mlp_ratio (int): ratio of mlp hidden dim to embedding dim for MLP of transformer layers
            qkv_bias (bool): enable bias for qkv if True
            drop_rate (float): dropout rate for MLP of transformer layers, MSA final projection layer, and classifier
            attn_drop_rate (float): attention dropout rate
            drop_path_rate (float): stochastic depth rate
            norm_layer: (nn.Module): normalization layer for transformer layers
            act_layer: (nn.Module): activation layer in MLP of transformer layers
            pad_type: str: Type of padding to use '' for PyTorch symmetric, 'same' for TF SAME
            weight_init: (str): weight init scheme
            global_pool: (str): type of pooling operation to apply to final feature map

        Notes:
            - Default values follow NesT-B from the original Jax code.
            - `embed_dims`, `num_heads`, `depths` should be ints or tuples with length `num_levels`.
            - For those following the paper, Table A1 may have errors!
                - https://github.com/google-research/nested-transformer/issues/2
        
embed_dimsr#   depthszRequire `len(z) == num_levels`r6   r   r   z Model only handles square inputsz*`patch_size` must divide `img_size` evenlyr2   zUFirst level blocks don't fit evenly. Check `img_size`, `patch_size`, and `num_levels`F)img_size
patch_sizein_chansrx   flattenc                 S   s   g | ]}|  qS r0   )tolist)ru   r=   r0   r0   r1   ry   [      z!Nest.__init__.<locals>.<listcomp>N)rZ   r-   r*   r(   rV   r[   rS   rg   zlevels.)Znum_chsZ	reductionmoduler3   Z	pool_type)/r!   r"   locals
isinstancecollectionsabcSequencer~   r   num_classesnum_featuresZfeature_infor   r   r^   	drop_rate
num_levelsr   rK   ZarangeZflipr   r   rm   rn   rX   ri   r   patch_embedZnum_patchesr   Zlinspacesumsplitr   appendrs   dictr   levelsre   r
   global_poolr'   	head_dropheadinit_weights)r+   r   r   r   r   r   r#   r   r   rZ   r-   r   Zproj_drop_rateZattn_drop_rateZdrop_path_rater[   rS   rg   Zweight_initr   
param_nameZparam_valuer   Zdp_ratesZprev_dimZcurr_striderv   r,   r   r.   r0   r1   r"      s    1



 
" zNest.__init__c                 C   sZ   |dv sJ d|v r"t | j nd}| jD ]}t|jdddd q,ttt|d|  d S )	N)nlhbr`   r   r   {Gz?r5   r3   stdab)	head_bias)	rm   logr   r   r   r}   r   r   _init_nest_weights)r+   moder   levelr0   r0   r1   r     s
    
zNest.init_weightsc                 C   s   dd t t| jD S )Nc                 S   s   h | ]}d | dqS )zlevel.z
.pos_embedr0   rt   r0   r0   r1   	<setcomp>  r   z'Nest.no_weight_decay.<locals>.<setcomp>)r   r~   r   r+   r0   r0   r1   no_weight_decay  s    zNest.no_weight_decayFc                 C   s"   t d|rdndd fddgd}|S )Nz^patch_embedz^levels\.(\d+)z*^levels\.(\d+)\.transformer_encoder\.(\d+))z"^levels\.(\d+)\.(?:pool|pos_embed))r   )z^norm)i )stemblocks)r   )r+   ZcoarseZmatcherr0   r0   r1   group_matcher  s    zNest.group_matcherc                 C   s   | j D ]
}||_qd S r\   )r   rz   )r+   enablelr0   r0   r1   set_grad_checkpointing  s    
zNest.set_grad_checkpointingc                 C   s   | j S r\   )r   r   r0   r0   r1   get_classifier  s    zNest.get_classifierc                 C   s$   || _ t| j| j |d\| _| _d S )Nr   )r   r
   r   r   r   )r+   r   r   r0   r0   r1   reset_classifier  s    
zNest.reset_classifierc                 C   s:   |  |}| |}| |dddddddd}|S )Nr   r3   r   r   )r   r   re   r9   rh   r0   r0   r1   forward_features  s    

"zNest.forward_features)
pre_logitsc                 C   s&   |  |}| |}|r|S | |S r\   )r   r   r   )r+   r=   r   r0   r0   r1   forward_head  s    

zNest.forward_headc                 C   s   |  |}| |}|S r\   )r   r   rh   r0   r0   r1   rF     s    

zNest.forward)r   r   r2   r   r   r   r   r   rR   Tr   r   r   r   NNr`   r`   r   )r`   )F)T)r   )F)rG   rH   rI   rJ   r"   rK   rL   ignorer   r   r   r   r   r   r   rN   r   rF   rP   r0   r0   r.   r1   r      sH                       


r`   r   r   namer   c                 C   s   t | tjrf|dr:t| jdddd tj| j| qt| jdddd | jdurtj	| j n6t | tj
rt| jdddd | jdurtj	| j dS )zn NesT weight initialization
    Can replicate Jax implementation. Otherwise follows vision_transformer.py
    r   r   r5   r3   r   N)r   r   r%   
startswithr   ZweightinitZ	constant_r    Zzeros_ZConv2dr   r0   r0   r1   r     s    


r   c                 C   s   t d| j|j | jd }|jdd \}}tt|| }t| tt|dddd} tj	| ||gddd} t
| ddddtt|} | S )	z
    Rescale the grid of position embeddings when loading from state_dict
    Expected shape of position embeddings is (1, T, N, C), and considers only square images
    z$Resized position embedding: %s to %sr3   r   r   r   bicubicF)sizer   Zalign_corners)_loggerinfor7   rX   rm   rn   rr   r9   r:   Zinterpolaterl   )ZposembZ
posemb_newZseq_length_oldZnum_blocks_newZseq_length_newZsize_newr0   r0   r1   resize_pos_embed  s    
  r   c                 C   sN   dd |   D }|D ]2}| | jt||jkrt| | t||| |< q| S )z4 resize positional embeddings of pretrained weights c                 S   s   g | ]}| d r|qS )Z
pos_embed_)r   )ru   rC   r0   r0   r1   ry     r   z(checkpoint_filter_fn.<locals>.<listcomp>)keysr7   getattrr   )Z
state_dictmodelZpos_embed_keysrC   r0   r0   r1   checkpoint_filter_fn  s
    r   Fc                 K   s&   t t| |ftdddtd|}|S )N)r   r   r3   T)Zout_indicesZflatten_sequential)Zfeature_cfgZpretrained_filter_fn)r   r   r   r   )variant
pretrainedkwargsr   r0   r0   r1   _create_nest  s    
	r   c                 K   s$   | ddddgdddt tddd	|S )
Nr   )r   r   r      g      ?r   Tzpatch_embed.projr   )urlr   Z
input_sizeZ	pool_sizeZcrop_pctinterpolationZfixed_input_sizemeanr   Z
first_conv
classifierr   )r   r   r0   r0   r1   _cfg  s    
r   ztimm/)Z	hf_hub_id)znest_base.untrainedznest_small.untrainedznest_tiny.untrainedznest_base_jx.goog_in1kznest_small_jx.goog_in1kznest_tiny_jx.goog_in1k)returnc                 K   s,   t f dddd|}tdd| i|}|S ) Nest-B @ 224x224
    r   r   r   r   	nest_baser   )r   r   r   r   r   Zmodel_kwargsr   r0   r0   r1   r     s    r   c                 K   s,   t f dddd|}tdd| i|}|S ) Nest-S @ 224x224
    `      i  r         r   r   
nest_smallr   )r   r   r   r0   r0   r1   r     s    r   c                 K   s,   t f dddd|}tdd| i|}|S ) Nest-T @ 224x224
    r   r   r3   r3   r   r   	nest_tinyr   )r   r   r   r0   r0   r1   r     s    r   c                 K   s8   | dd tf dddd|}td	d| i|}|S )
r   rg   samer   r   r   r   nest_base_jxr   )r   
setdefaultr   r   r   r0   r0   r1   r   $  s    r   c                 K   s8   | dd tf dddd|}td	d| i|}|S )
r   rg   r   r   r   r   r   nest_small_jxr   )r   r   r   r0   r0   r1   r   /  s    r   c                 K   s8   | dd tf dddd|}td	d| i|}|S )
r   rg   r   r   r   r   r   nest_tiny_jxr   )r   r   r   r0   r0   r1   r   9  s    r   )Zjx_nest_baseZjx_nest_smallZjx_nest_tiny)r`   r   )F)r`   )F)F)F)F)F)F)ArJ   collections.abcr   loggingrm   	functoolsr   rK   Ztorch.nn.functionalr   Z
functionalr:   Z	timm.datar   r   Ztimm.layersr   r   r	   r
   r   r   r   r   r   r   r   Z_builderr   Z_features_fxr   Z_manipulater   r   	_registryr   r   r   __all__	getLoggerrG   r   Moduler   rQ   r_   rX   rl   rr   rs   r   strfloatr   r   r   r   r   Zdefault_cfgsr   r   r   r   r   r   r0   r0   r0   r1   <module>   sn    
,,B E	

	
		