a
    d|                     @   sT  d Z ddlZddlZddlmZmZmZmZ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 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dl$m%Z% dgZ&e'e(Z)ee*ee*e*f f Z+e*dddZ,ee*e*e*dddZ-e*e*dddZ.G dd dej/Z0G dd dej/Z1G dd dej/Z2G dd dej/Z3G dd dej/Z4dd  Z5d\d"d#Z6d]d%d&Z7e!e7d'd(d)e7d'd*d)e7d'd+d,d-d.d/e7d'd0d)e7d'd1d,d-d.d/e7d'd2d)e7d'd3d)e7d'd4d)e7d'd5d,d-d.d/e7d'd6d)e7d'd7d8d9e7d'd:d8d9e7d'd;d8d9e7d'd<d,d-d.d8d=e7d'd>d8d9e7d'd?d,d-d.d8d=e7d'd@d)e7d'dAd)e7d'dBd)dCZ8e"d^e4dDdEdFZ9e"d_e4dDdGdHZ:e"d`e4dDdIdJZ;e"dae4dDdKdLZ<e"dbe4dDdMdNZ=e"dce4dDdOdPZ>e"dde4dDdQdRZ?e"dee4dDdSdTZ@e"dfe4dDdUdVZAe#e(dWdXdYdZd[ dS )ga   Swin Transformer
A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows`
    - https://arxiv.org/pdf/2103.14030

Code/weights from https://github.com/microsoft/Swin-Transformer, original copyright/license info below

S3 (AutoFormerV2, https://arxiv.org/abs/2111.14725) Swin weights from
    - https://github.com/microsoft/Cream/tree/main/AutoFormerV2

Modifications and additions for timm hacked together by / Copyright 2021, Ross Wightman
    N)CallableListOptionalTupleUnionIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)	
PatchEmbedMlpDropPathClassifierHead	to_2tuple	to_ntupletrunc_normal__assertuse_fused_attn   )build_model_with_cfg)register_notrace_function)checkpoint_seqnamed_apply)generate_default_cfgsregister_modelregister_model_deprecations)get_init_weights_vitSwinTransformer)window_sizec                 C   sR   | j \}}}}| ||| ||| ||} | dddddd d|||}|S )z
    Args:
        x: (B, H, W, C)
        window_size (int): window size

    Returns:
        windows: (num_windows*B, window_size, window_size, C)
    r   r               shapeviewpermute
contiguous)xr   BHWCwindows r.   e/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/swin_transformer.pywindow_partition)   s    	$r0   )r   r*   r+   c                 C   sN   | j d }| d|| || |||}|dddddd d|||}|S )z
    Args:
        windows: (num_windows*B, window_size, window_size, C)
        window_size (int): Window size
        H (int): Height of image
        W (int): Width of image

    Returns:
        x: (B, H, W, C)
    r"   r   r   r   r   r    r!   r#   )r-   r   r*   r+   r,   r(   r.   r.   r/   window_reverse8   s    
$r1   )win_hwin_wc                 C   s   t t t | t |g}t |d}|d d d d d f |d d d d d f  }|ddd }|d d d d df  | d 7  < |d d d d df  |d 7  < |d d d d df  d| d 9  < |dS )Nr   r   r   r"   )torchstackZmeshgridZarangeflattenr&   r'   sum)r2   r3   ZcoordsZcoords_flattenZrelative_coordsr.   r.   r/   get_relative_position_indexJ   s     ,""&r8   c                	       sp   e Zd ZU dZejje ed< de	e	e
e	 eeeed fdd	Zejd
ddZde
ej dddZ  ZS )WindowAttentionz Window based multi-head self attention (W-MSA) module with relative position bias.
    It supports shifted and non-shifted windows.
    
fused_attnN   T        )dim	num_headshead_dimr   qkv_bias	attn_drop	proj_dropc                    s   t    || _t|| _| j\}}	||	 | _|| _|p>|| }|| }
|d | _tdd| _	t
td| d d|	 d  || _| dt||	 t
j||
d |d| _t
|| _t
|
|| _t
|| _t| jd	d
 t
jdd| _dS )a  
        Args:
            dim: Number of input channels.
            num_heads: Number of attention heads.
            head_dim: Number of channels per head (dim // num_heads if not set)
            window_size: The height and width of the window.
            qkv_bias:  If True, add a learnable bias to query, key, value.
            attn_drop: Dropout ratio of attention weight.
            proj_drop: Dropout ratio of output.
        g      T)Zexperimentalr   r   relative_position_indexr   Zbiasg{Gz?)stdr"   )r=   N)super__init__r=   r   r   window_arear>   scaler   r:   nn	Parameterr4   zerosrelative_position_bias_tableregister_bufferr8   LinearqkvZDropoutrA   projrB   r   ZSoftmaxsoftmax)selfr=   r>   r?   r   r@   rA   rB   r2   r3   Zattn_dim	__class__r.   r/   rG   \   s$    




(zWindowAttention.__init__returnc                 C   s<   | j | jd | j| jd}|ddd }|dS )Nr"   r   r   r   )rM   rC   r%   rH   r&   r'   	unsqueeze)rS   Zrelative_position_biasr.   r.   r/   _get_rel_pos_bias   s    

z!WindowAttention._get_rel_pos_biasmaskc                 C   sx  |j \}}}| |||d| jdddddd}|d\}}}	| jr|  }
|dur|j d }|d|d||	|| d| jdd}|
|d| j|| }
t
jjj|||	|
| jjd}n|| j }||d	d }||   }|dur.|j d }|d|| j|||dd }|d| j||}| |}| |}||	 }|dd||d}| |}| |}|S )
z
        Args:
            x: input features with shape of (num_windows*B, N, C)
            mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None
        r   r"   r   r   r   r    N)	attn_maskZ	dropout_p)r$   rP   reshaper>   r&   Zunbindr:   rY   r%   expandr4   rJ   Z
functionalZscaled_dot_product_attentionrA   prI   Z	transposerX   rR   rQ   rB   )rS   r(   r[   ZB_Nr,   rP   qkvr\   Znum_winattnr.   r.   r/   forward   s8    (
&


$



zWindowAttention.forward)Nr;   Tr<   r<   )N)__name__
__module____qualname____doc__r4   jitFinalbool__annotations__intr   _int_or_tuple_2_tfloatrG   ZTensorrY   rf   __classcell__r.   r.   rT   r/   r9   V   s$   
     -r9   c                       sf   e Zd ZdZdddddddddejejfeeee	e eee
ee
e
e
eed	 fd
dZdd Z  ZS )SwinTransformerBlockz Swin Transformer Block.
    r    Nr;   r         @Tr<   )r=   input_resolutionr>   r?   r   
shift_size	mlp_ratior@   rB   rA   	drop_path	act_layer
norm_layerc              	      s  t    || _|| _|| _|| _|| _t| j| jkrJd| _t| j| _d| j  krb| jk sln J d||| _t	|||t
| j||
|	d| _|dkrt|nt | _||| _t|t|| ||	d| _| jdkr| j\}}td||df}d}td| j t| j | j t| j dfD ]Z}td| j t| j | j t| j dfD ]&}||dd||ddf< |d7 }qZq*t|| j}|d| j| j }|d|d	 }||dktd
|dktd}nd}| d| dS )a  
        Args:
            dim: Number of input channels.
            input_resolution: Input resolution.
            window_size: Window size.
            num_heads: Number of attention heads.
            head_dim: Enforce the number of channels per head
            shift_size: Shift size for SW-MSA.
            mlp_ratio: Ratio of mlp hidden dim to embedding dim.
            qkv_bias: If True, add a learnable bias to query, key, value.
            proj_drop: Dropout rate.
            attn_drop: Attention dropout rate.
            drop_path: Stochastic depth rate.
            act_layer: Activation layer.
            norm_layer: Normalization layer.
        r   z shift_size must in 0-window_size)r>   r?   r   r@   rA   rB   r<   )Zin_featuresZhidden_featuresry   Zdropr   Nr"   r   g      Yr\   )rF   rG   r=   ru   r   rv   rw   minnorm1r9   r   re   r   rJ   Identityrx   norm2r   ro   mlpr4   rL   slicer0   r%   rX   Zmasked_fillrq   rN   )rS   r=   ru   r>   r?   r   rv   rw   r@   rB   rA   rx   ry   rz   r*   r+   Zimg_maskZcnthwZmask_windowsr\   rT   r.   r/   rG      s`     
"




&zSwinTransformerBlock.__init__c           
      C   s8  |j \}}}}t|| jd kd t|| jd kd |}| |}| jdkrltj|| j | j fdd}n|}t|| j}|	d| j| j |}| j
|| jd}	|		d| j| j|}	t|	| j||}| jdkrtj|| j| jfdd}n|}|| | }||d|}|| | | | }|||||}|S )Nr   zinput feature has wrong sizer   )r   r   )Zshiftsdimsr"   rZ   )r$   r   ru   r|   rv   r4   Zrollr0   r   r%   re   r\   r1   rx   r^   r   r~   )
rS   r(   r)   r*   r+   r,   ZshortcutZ	shifted_xZ	x_windowsZattn_windowsr.   r.   r/   rf     s*    


zSwinTransformerBlock.forward)rg   rh   ri   rj   rJ   ZGELU	LayerNormro   rp   r   rq   rm   r   rG   rf   rr   r.   r.   rT   r/   rs      s8   Wrs   c                       s>   e Zd ZdZdejfeee ed fddZ	dd Z
  ZS )PatchMergingz Patch Merging Layer.
    Nr=   out_dimrz   c                    sH   t    || _|pd| | _|d| | _tjd| | jdd| _dS )z
        Args:
            dim: Number of input channels.
            out_dim: Number of output channels (or 2 * dim if None)
            norm_layer: Normalization layer.
        r   r    FrD   N)rF   rG   r=   r   normrJ   rO   	reduction)rS   r=   r   rz   rT   r.   r/   rG   =  s
    
zPatchMerging.__init__c                 C   s   |j \}}}}t|d dkd| d t|d dkd| d |||d d|d d|dddddd	d}| |}| |}|S )
Nr   r   z
x height (z) is not even.z	x width (r   r   r    r!   )r$   r   r^   r&   r6   r   r   )rS   r(   r)   r*   r+   r,   r.   r.   r/   rf   O  s    2

zPatchMerging.forward)rg   rh   ri   rj   rJ   r   ro   r   r   rG   rf   rr   r.   r.   rT   r/   r   9  s   r   c                       sx   e Zd ZdZdddddddddejf
eeeeef eeee	e e
eeeeeee ef ed fd	d
Zdd Z  ZS )SwinTransformerStagez3 A basic Swin Transformer layer for one stage.
    Tr    Nr;   rt   r<   r=   r   ru   depth
downsampler>   r?   r   rw   r@   rB   rA   rx   rz   c                    s   t    |	_|	_|r,tdd |D n|	_|	_d	_|rTt|d	_	n|ks`J t
 	_	t
j 	
fddt|D  	_dS )a  
        Args:
            dim: Number of input channels.
            input_resolution: Input resolution.
            depth: Number of blocks.
            downsample: Downsample layer at the end of the layer.
            num_heads: Number of attention heads.
            head_dim: Channels per head (dim // num_heads if not set)
            window_size: Local window size.
            mlp_ratio: Ratio of mlp hidden dim to embedding dim.
            qkv_bias: If True, add a learnable bias to query, key, value.
            proj_drop: Projection dropout rate.
            attn_drop: Attention dropout rate.
            drop_path: Stochastic depth rate.
            norm_layer: Normalization layer.
        c                 s   s   | ]}|d  V  qdS )r   Nr.   .0ir.   r.   r/   	<genexpr>      z0SwinTransformerStage.__init__.<locals>.<genexpr>Fr   c                    sT   g | ]L}t 	j
|d  dkr&dn
d   ttrF| ndqS )r   r   )r=   ru   r>   r?   r   rv   rw   r@   rB   rA   rx   rz   )rs   output_resolution
isinstancelistr   rA   rx   r?   rw   rz   r>   r   rB   r@   rS   r   r.   r/   
<listcomp>  s   z1SwinTransformerStage.__init__.<locals>.<listcomp>N)rF   rG   r=   ru   tupler   r   grad_checkpointingr   r   rJ   r}   
Sequentialrangeblocks)rS   r=   r   ru   r   r   r>   r?   r   rw   r@   rB   rA   rx   rz   rT   r   r/   rG   ]  s"    !


"zSwinTransformerStage.__init__c                 C   s6   |  |}| jr(tj s(t| j|}n
| |}|S N)r   r   r4   rk   Zis_scriptingr   r   rS   r(   r.   r.   r/   rf     s
    

zSwinTransformerStage.forward)rg   rh   ri   rj   rJ   r   ro   r   rm   r   rp   rq   r   r   r   rG   rf   rr   r.   r.   rT   r/   r   Y  s8   	
Er   c                       s  e Zd ZdZdddddddd	d
dddddddejdfeeeeeee	edf e	edf e
e eeeeeeeeeef e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 Swin Transformer

    A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows`  -
          https://arxiv.org/pdf/2103.14030
       r    r     avg`   r   r      r   r   r         Nr;   rt   Tr<   g? .)img_size
patch_sizein_chansnum_classesglobal_pool	embed_dimdepthsr>   r?   r   rw   r@   	drop_rateproj_drop_rateattn_drop_ratedrop_path_raterz   weight_initc                    s  t    |dv sJ || _|| _d| _t|| _ | _t d| jd   | _	g | _
t ttfs| fddt| jD  t||| d |dd| _| jj| _t| j|	}	t| j|
}
t| j|}d	d td|t||D }g } d }d}t| jD ]} | }|t||| jd | | jd | f|| |dk|| |	| |
| || ||||| |d
g7 }|}|dkr|d9 }|  j
t|d| d| dg7  _
qtj| | _|| j	| _t| j	|||| jd| _|dkr|  | dS )aN  
        Args:
            img_size: Input image size.
            patch_size: Patch size.
            in_chans: Number of input image channels.
            num_classes: Number of classes for classification head.
            embed_dim: Patch embedding dimension.
            depths: Depth of each Swin Transformer layer.
            num_heads: Number of attention heads in different layers.
            head_dim: Dimension of self-attention heads.
            window_size: Window size.
            mlp_ratio: Ratio of mlp hidden dim to embedding dim.
            qkv_bias: If True, add a learnable bias to query, key, value.
            drop_rate: Dropout rate.
            attn_drop_rate (float): Attention dropout rate.
            drop_path_rate (float): Stochastic depth rate.
            norm_layer (nn.Module): Normalization layer.
        )r   r   ZNHWCr   r   c                    s   g | ]}t  d |  qS )r   )ro   r   r   r.   r/   r     r   z,SwinTransformer.__init__.<locals>.<listcomp>r   )r   r   r   r   rz   
output_fmtc                 S   s   g | ]}|  qS r.   )tolist)r   r(   r.   r.   r/   r     r   r   r    layers.)Znum_chsr   module)	pool_typer   Z	input_fmtskipN)!rF   rG   r   r   r   lenZ
num_layersr   ro   Znum_featuresZfeature_infor   r   r   r   r
   patch_embedZ	grid_sizeZ
patch_gridr   r4   Zlinspacer7   splitr   dictrJ   r   layersr   r   headinit_weights)rS   r   r   r   r   r   r   r   r>   r?   r   rw   r@   r   r   r   r   rz   r   kwargsZdprr   Zin_dimrI   r   r   rT   r   r/   rG     s|    (


"

(
zSwinTransformer.__init__c                 C   s<   |dv sJ d|v r"t | j nd}tt||d|  d S )N)ZjaxZjax_nlhbZmocor   Znlhbr<   )	head_bias)mathlogr   r   r   )rS   moder   r.   r.   r/   r   "  s    zSwinTransformer.init_weightsc                 C   s.   t  }|  D ]\}}d|v r|| q|S )NrM   )setZnamed_parametersadd)rS   Znwdn_r.   r.   r/   no_weight_decay(  s
    zSwinTransformer.no_weight_decayFc                 C   s   t d|rdng ddS )Nz^patch_embedz^layers\.(\d+)))z^layers\.(\d+).downsample)r   )z^layers\.(\d+)\.\w+\.(\d+)N)z^norm)i )stemr   )r   )rS   Zcoarser.   r.   r/   group_matcher0  s    zSwinTransformer.group_matcherc                 C   s   | j D ]
}||_qd S r   )r   r   )rS   enablelr.   r.   r/   set_grad_checkpointing;  s    
z&SwinTransformer.set_grad_checkpointingc                 C   s   | j jS r   )r   Zfc)rS   r.   r.   r/   get_classifier@  s    zSwinTransformer.get_classifierc                 C   s   || _ | jj||d d S )N)r   )r   r   reset)rS   r   r   r.   r.   r/   reset_classifierD  s    z SwinTransformer.reset_classifierc                 C   s"   |  |}| |}| |}|S r   )r   r   r   r   r.   r.   r/   forward_featuresH  s    


z SwinTransformer.forward_features
pre_logitsc                 C   s   |r| j |ddS |  |S )NTr   )r   )rS   r(   r   r.   r.   r/   forward_headN  s    zSwinTransformer.forward_headc                 C   s   |  |}| |}|S r   )r   r   r   r.   r.   r/   rf   Q  s    

zSwinTransformer.forward)r   )F)T)N)F)rg   rh   ri   rj   rJ   r   rp   ro   strr   r   rq   rm   r   r   rG   r4   rk   ignorer   r   r   r   r   r   r   r   rf   rr   r.   r.   rT   r/   r     sj   


o



c                 C   sl   d| v r| S ddl }i }| d| } | d| } |  D ].\}}|ddd |}|d	d
}|||< q8|S )zJ convert patch embedding weight from manual patchify + linear proj to convzhead.fc.weightr   Nmodel
state_dictzlayers.(\d+).downsamplec                 S   s   dt | dd  dS )Nr   r   z.downsample)ro   group)r(   r.   r.   r/   <lambda>`  r   z&checkpoint_filter_fn.<locals>.<lambda>zhead.zhead.fc.)regetitemssubreplace)r   r   r   Zout_dictrc   rd   r.   r.   r/   checkpoint_filter_fnW  s    
r   Fc                 K   sP   t dd t|ddD }|d|}tt| |fttd|dd|}|S )	Nc                 s   s   | ]\}}|V  qd S r   r.   )r   r   r   r.   r.   r/   r   g  r   z+_create_swin_transformer.<locals>.<genexpr>r   )r   r   r   r   out_indicesT)Zflatten_sequentialr   )Zpretrained_filter_fnZfeature_cfg)r   	enumerater   popr   r   r   r   )variant
pretrainedr   Zdefault_out_indicesr   r   r.   r.   r/   _create_swin_transformerf  s    
r   r   c                 K   s"   | ddddddt tddd	d
|S )Nr   )r   r   r   )r;   r;   g?ZbicubicTzpatch_embed.projzhead.fcZmit)urlr   
input_size	pool_sizecrop_pctinterpolationZfixed_input_sizemeanrE   Z
first_conv
classifierlicenser   )r   r   r.   r.   r/   _cfgs  s    r   ztimm/zvhttps://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_small_patch4_window7_224_22kto1k_finetune.pth)	hf_hub_idr   zlhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window7_224_22kto1k.pthzmhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window12_384_22kto1k.pth)r     r   )r   r   g      ?)r   r   r   r   r   zmhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window7_224_22kto1k.pthznhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window12_384_22kto1k.pthzdhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_tiny_patch4_window7_224.pthzehttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_small_patch4_window7_224.pthzdhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window7_224.pthzehttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window12_384.pthzuhttps://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_tiny_patch4_window7_224_22kto1k_finetune.pthzhhttps://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_tiny_patch4_window7_224_22k.pthiQU  )r   r   r   zihttps://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_small_patch4_window7_224_22k.pthzhhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window7_224_22k.pthzihttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window12_384_22k.pth)r   r   r   r   r   r   zihttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window7_224_22k.pthzjhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window12_384_22k.pthzbhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/s3_t-1d53f6a8.pthzbhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/s3_s-3bb4c69d.pthzbhttps://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/s3_b-a1e95db4.pth)z.swin_small_patch4_window7_224.ms_in22k_ft_in1kz-swin_base_patch4_window7_224.ms_in22k_ft_in1kz.swin_base_patch4_window12_384.ms_in22k_ft_in1kz.swin_large_patch4_window7_224.ms_in22k_ft_in1kz/swin_large_patch4_window12_384.ms_in22k_ft_in1kz$swin_tiny_patch4_window7_224.ms_in1kz%swin_small_patch4_window7_224.ms_in1kz$swin_base_patch4_window7_224.ms_in1kz%swin_base_patch4_window12_384.ms_in1kz-swin_tiny_patch4_window7_224.ms_in22k_ft_in1kz%swin_tiny_patch4_window7_224.ms_in22kz&swin_small_patch4_window7_224.ms_in22k%swin_base_patch4_window7_224.ms_in22k&swin_base_patch4_window12_384.ms_in22k&swin_large_patch4_window7_224.ms_in22k'swin_large_patch4_window12_384.ms_in22kzswin_s3_tiny_224.ms_in1kzswin_s3_small_224.ms_in1kzswin_s3_base_224.ms_in1krV   c                 K   s0   t dddddd}td	d| it |fi |S )
z+ Swin-T @ 224x224, trained ImageNet-1k
    r    r;   r   r   r   r   r   r   r   r>   swin_tiny_patch4_window7_224r   )r   r   r   r   r   Z
model_argsr.   r.   r/   r     s     r   c                 K   s0   t dddddd}td	d| it |fi |S )
z Swin-S @ 224x224
    r    r;   r   r   r      r   r   r   swin_small_patch4_window7_224r   )r   r   r   r.   r.   r/   r     s     r   c                 K   s0   t dddddd}td	d| it |fi |S )
z Swin-B @ 224x224
    r    r;      r   r              r   swin_base_patch4_window7_224r   )r  r   r   r.   r.   r/   r    s     r  c                 K   s0   t dddddd}td	d| it |fi |S )
z Swin-B @ 384x384
    r    r   r  r   r  r   swin_base_patch4_window12_384r   )r  r   r   r.   r.   r/   r    s     r  c                 K   s0   t dddddd}td	d| it |fi |S )
z Swin-L @ 224x224
    r    r;      r   r   r   r   0   r   swin_large_patch4_window7_224r   )r  r   r   r.   r.   r/   r    s     r  c                 K   s0   t dddddd}td	d| it |fi |S )
z Swin-L @ 384x384
    r    r   r  r   r	  r   swin_large_patch4_window12_384r   )r  r   r   r.   r.   r/   r    s     r  c                 K   s0   t dddddd}td	d| it |fi |S )
z; Swin-S3-T @ 224x224, https://arxiv.org/abs/2111.14725
    r    r;   r;      r;   r   r   r   r   swin_s3_tiny_224r   )r  r   r   r.   r.   r/   r    s    
r  c                 K   s0   t dddddd}td	d| it |fi |S )
z; Swin-S3-S @ 224x224, https://arxiv.org/abs/2111.14725
    r    )r  r  r  r;   r   r   r   r   swin_s3_small_224r   )r  r   r   r.   r.   r/   r    s    
r  c                 K   s0   t dddddd}td	d| it |fi |S )
z; Swin-S3-B @ 224x224, https://arxiv.org/abs/2111.14725
    r    r  r   )r   r      r   r   r   swin_s3_base_224r   )r  r   r   r.   r.   r/   r    s    
r  r   r   r   r   )Z"swin_base_patch4_window7_224_in22kZ#swin_base_patch4_window12_384_in22kZ#swin_large_patch4_window7_224_in22kZ$swin_large_patch4_window12_384_in22k)F)r   )F)F)F)F)F)F)F)F)F)Brj   loggingr   typingr   r   r   r   r   r4   Ztorch.nnrJ   Z	timm.datar   r	   Ztimm.layersr
   r   r   r   r   r   r   r   r   Z_builderr   Z_features_fxr   Z_manipulater   r   	_registryr   r   r   Zvision_transformerr   __all__	getLoggerrg   Z_loggerro   rp   r0   r1   r8   Moduler9   rs   r   r   r   r   r   r   Zdefault_cfgsr   r   r  r  r  r  r  r  r  r.   r.   r.   r/   <module>   s  ,
`  S ,

K