
    ^j&                        d Z ddlZddl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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 ddl 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j`                  e1      Z2ee3e
e3e3f   f   Z4dejj                  de
e3e3f   dejj                  fdZ6e%dejj                  de
e3e3f   de3de3dejj                  f
d       Z7dbde3de3dejj                  fdZ8 G d dejr                        Z: G d dejr                        Z; G d dejr                        Z< G d d ejr                        Z= G d! dejr                        Z>d"e?d#ejr                  dee@ejj                  f   fd$ZAdcd%e@d&eBde>fd'ZCddd(e@dee@ef   fd)ZD e*i d* eDd+d,-      d. eDd+d/-      d0 eDd+d1d2d3d45      d6 eDd+d7-      d8 eDd+d9d2d3d45      d: eDd+d;-      d< eDd+d=-      d> eDd+d?-      d@ eDd+dAd2d3d45      dB eDd+dC-      dD eDd+dEdFG      dH eDd+dIdFG      dJ eDd+dKdFG      dL eDd+dMd2d3d4dFN      dO eDd+dPdFG      dQ eDd+dRd2d3d4dFN      dS eDd+dT-       eDd+dU-       eDd+dV-      dW      ZEe+dcde>fdX       ZFe+dcde>fdY       ZGe+dcde>fdZ       ZHe+dcde>fd[       ZIe+dcde>fd\       ZJe+dcde>fd]       ZKe+dcde>fd^       ZLe+dcde>fd_       ZMe+dcde>fd`       ZN e,e1dJdLdOdQda       y)ea   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)	AnyDictCallableListOptionalSetTupleUnionTypeIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)
PatchEmbedMlpDropPathcalculate_drop_path_ratesClassifierHead	to_2tuple	to_ntupletrunc_normal_use_fused_attnresize_rel_pos_bias_tableresample_patch_embedndgrid   )build_model_with_cfg)feature_take_indices)register_notrace_function)checkpoint_seqnamed_apply)generate_default_cfgsregister_modelregister_model_deprecations)get_init_weights_vitSwinTransformerxwindow_sizereturnc                     | j                   \  }}}}| j                  |||d   z  |d   ||d   z  |d   |      } | j                  dddddd      j                         j                  d|d   |d   |      }|S )zPartition into non-overlapping windows.

    Args:
        x: Input tokens with shape [B, H, W, C].
        window_size: Window size.

    Returns:
        Windows after partition with shape [B * num_windows, window_size, window_size, C].
    r   r               shapeviewpermute
contiguous)r&   r'   BHWCwindowss          g/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/swin_transformer.pywindow_partitionr:   *   s     JAq!Q	q!{1~%{1~qKN7JKXYN\]^Aii1aAq)446;;BAP[\]P^`abGN    r8   r5   r6   c                     | j                   d   }| j                  d||d   z  ||d   z  |d   |d   |      }|j                  dddddd      j                         j                  d|||      }|S )a
  Reverse window partition.

    Args:
        windows: Windows with shape (num_windows*B, window_size, window_size, C).
        window_size: Window size.
        H: Height of image.
        W: Width of image.

    Returns:
        Tensor with shape (B, H, W, C).
    r.   r   r   r*   r+   r,   r-   r/   )r8   r'   r5   r6   r7   r&   s         r9   window_reverser=   =   s     	bARk!n,a;q>.A;q>S^_`SacdeA			!Q1a#..055b!QBAHr;   win_hwin_wc           
      "   t        j                  t        t        j                  | |t         j                        t        j                  ||t         j                                    }t        j
                  |d      }|dddddf   |dddddf   z
  }|j                  ddd      j                         }|dddddfxx   | dz
  z  cc<   |dddddfxx   |dz
  z  cc<   |dddddfxx   d|z  dz
  z  cc<   |j                  d      S )zGet pair-wise relative position index for each token inside the window.

    Args:
        win_h: Window height.
        win_w: Window width.

    Returns:
        Relative position index tensor.
    devicedtyper   Nr+   r   r.   )	torchstackr   arangelongflattenr2   r3   sum)r>   r?   rB   coordscoords_flattenrelative_coordss         r9   get_relative_position_indexrM   P   s     [[U6<U6< F ]]61-N$Q4Z0>!T1*3MMO%--aA6AACOAq!G	)Aq!G	)Aq!GE	A-r""r;   c                   :    e Zd ZU dZej
                  j                  e   ed<   	 	 	 	 	 	 	 dde	de	de
e	   deded	ed
ef fdZddZddZdee	e	f   ddfdZdej$                  fdZddej$                  de
ej$                     dej$                  fdZddZ xZS )WindowAttentionzWindow based multi-head self attention (W-MSA) module with relative position bias.

    Supports both shifted and non-shifted windows.
    
fused_attnNdim	num_headshead_dimr'   qkv_bias	attn_drop	proj_dropc
           	         ||	d}
t         |           || _        t        |      | _        | j                  \  }}||z  | _        || _        |xs ||z  }||z  }|dz  | _        t        d      | _	        t        j                  t        j                  d|z  dz
  d|z  dz
  z  |fi |
      | _        | j                  dt        j                  ||z  ||z  |t        j                         d	       t        j"                  ||d
z  fd|i|
| _        t        j&                  |      | _        t        j"                  ||fi |
| _        t        j&                  |      | _        t        j.                  d      | _        | j3                          y)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.
        rA   g      T)experimentalr+   r   relative_position_indexF
persistentr*   biasr.   )rQ   N)super__init__rQ   r   r'   window_arearR   scaler   rP   nn	ParameterrD   emptyrelative_position_bias_tableregister_bufferrG   LinearqkvDropoutrU   projrV   Softmaxsoftmaxreset_parameters)selfrQ   rR   rS   r'   rT   rU   rV   rB   rC   ddr>   r?   attn_dim	__class__s                 r9   r^   zWindowAttention.__init__o   sm   , /$[1''u 5="/si/i'%
(d; -/LLKKUQ1u9q=99KK-M) 	%KKuu}V5::V 	 	
 99S(Q,DXDDI.IIh2r2	I.zzb) 	r;   r(   c                 R    t        | j                  d       | j                          y)"Initialize parameters and buffers.g{Gz?)stdN)r   rd   _init_buffersrm   s    r9   rl   z WindowAttention.reset_parameters   s    d77SAr;   c                     | j                   \  }}| j                  j                  t        ||| j                  j                               y).Compute and fill non-persistent buffer values.rB   N)r'   rY   copy_rM   rB   )rm   r>   r?   s      r9   rt   zWindowAttention._init_buffers   s=    ''u$$**'uT=Y=Y=`=`a	
r;   c           	         t        |      }|| j                  k(  ry|| _        | j                  \  }}||z  | _        t        j                         5  d|z  dz
  d|z  dz
  z  | j
                  f}t        j                  t        | j                  | j                  |            | _	        | j                  dt        ||| j                  j                        d       ddd       y# 1 sw Y   yxY w)	zzUpdate window size & interpolate position embeddings
        Args:
            window_size (int): New window size
        Nr+   r   new_window_sizenew_bias_shaperY   rx   FrZ   )r   r'   r_   rD   no_gradrR   ra   rb   r   rd   re   rM   rB   )rm   r'   r>   r?   r}   s        r9   set_window_sizezWindowAttention.set_window_size   s    
  ,$***&''u 5=]]_ 	%i!mE	A>NN02)55$($4$4#11D-   )+E5AbAbAiAij  ! 	 	 	s   BC%%C.c                     | j                   | j                  j                  d         j                  | j                  | j                  d      }|j	                  ddd      j                         }|j                  d      S )Nr.   r+   r   r   )rd   rY   r1   r_   r2   r3   	unsqueeze)rm   relative_position_biass     r9   _get_rel_pos_biasz!WindowAttention._get_rel_pos_bias   ss    !%!B!B((--b1"33748H8H$JZJZ\^3_ 	!7!?!?1a!H!S!S!U%//22r;   r&   maskc                    |j                   \  }}}| j                  |      j                  ||d| j                  d      j	                  ddddd      }|j                  d      \  }}}	| j                  r| j                         }
|e|j                   d   }|j                  d|d||      j                  ||z  d| j                  dd      }|
|j                  d| j                  ||      z   }
t        j                  j                  j                  |||	|
| j                  r| j                  j                   nd      }n|| j"                  z  }||j%                  d	d      z  }|| j                         z   }|m|j                   d   }|j                  d|| j                  ||      |j'                  d      j'                  d      z   }|j                  d| j                  ||      }| j)                  |      }| j                  |      }||	z  }|j%                  dd      j                  ||d      }| j+                  |      }| j-                  |      }|S )
a  Forward pass.

        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.

        Returns:
            Output features with shape of (num_windows*B, N, C).
        r*   r.   r+   r   r   r,           )	attn_mask	dropout_p)r0   rg   reshaperR   r2   unbindrP   r   r1   expandrD   ra   
functionalscaled_dot_product_attentiontrainingrU   pr`   	transposer   rk   ri   rV   )rm   r&   r   B_Nr7   rg   qkvr   num_winattns                r9   forwardzWindowAttention.forward   s    77Aqhhqk!!"aDNNB?GG1aQRTUV**Q-1a??..0I**Q-yyGQ15<<R7]BPTP^P^`bdfg%RA(NN	##@@1a#.2mm$..** A A DJJAq{{2r**D$0022D**Q-yyWdnnaCdnnUVFWFaFabcFddyyT^^Q:<<%D>>$'DqAKK1%%b!R0IIaLNN1r;   c                 $    | j                          yz"Initialize non-persistent buffers.Nrt   ru   s    r9   init_non_persistent_buffersz+WindowAttention.init_non_persistent_buffers       r;   )N   Tr   r   NNr(   NN)__name__
__module____qualname____doc__rD   jitFinalbool__annotations__intr   _int_or_tuple_2_tfloatr^   rl   rt   r	   r   Tensorr   r   r   __classcell__rp   s   @r9   rO   rO   h   s     		%% '+-.!!!4 4  4  sm	4 
 +4  4  4  4 l

5c? t 235<< 3( (Xell-C (u|| (Tr;   rO   c                    V    e Zd ZdZddddddddd	d	d	ej
                  ej                  ddfd
ededede	e   dedede
de
dede
dedededeej                     deej                     f fdZd)dZd)dZ	 	 	 d*de	ej$                     de	ej&                     de	ej(                     de	ej$                     fd Z	 d+d!eeeeef   f   d"e	eeeeef   f      deeeef   eeef   f   fd#Z	 d+d$eeef   deeef   de	e
   fd%Zd& Zdej$                  dej$                  fd'Zd)d(Z xZS ),SwinTransformerBlockzkSwin Transformer Block.

    A transformer block with window-based self-attention and shifted windows.
    r,   Nr   r   F      @Tr   rQ   input_resolutionrR   rS   r'   
shift_sizealways_partitiondynamic_mask	mlp_ratiorT   rV   rU   	drop_path	act_layer
norm_layerc           
         ||d}t         |           || _        || _        t	        |      | _        || _        || _        | j                  ||      \  | _	        | _
        | j                  d   | j                  d   z  | _        |	| _         ||fi || _        t        |f||| j                  |
||d|| _        |dkD  rt!        |      nt#        j$                         | _         ||fi || _        t+        d|t-        ||	z        ||d|| _        |dkD  rt!        |      nt#        j$                         | _        | j3                  ddd	
       | j5                          y)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.
            always_partition: Always partition into full windows and shift
            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.
        rA   r   r   )rR   rS   r'   rT   rU   rV   r   )in_featureshidden_featuresr   dropr   NFrZ    )r]   r^   rQ   r   r   target_shift_sizer   r   _calc_window_shiftr'   r   r_   r   norm1rO   r   r   ra   Identity
drop_path1norm2r   r   mlp
drop_path2re   rl   )rm   rQ   r   rR   rS   r'   r   r   r   r   rT   rV   rU   r   r   r   rB   rC   rn   rp   s                      r9   r^   zSwinTransformerBlock.__init__  ss   J / 0!*:!6 0(,0,C,CKQ[,\)$/++A.1A1A!1DD"*r*
#	
((	
 	
	 2;R(9-R[[]*r*
 
i0	

 
 2;R(9-R[[] 	[$5A 	r;   r(   c                 $    | j                          y)rr   Nr   ru   s    r9   rl   z%SwinTransformerBlock.reset_parametersR  r   r;   c                     | j                   sh| j                  j                  j                  }| j                  j                  j                  }| j                  ||      }| j                  d|d       yy)rw   rA   r   FrZ   N)r   r   weightrB   rC   get_attn_maskre   )rm   rB   rC   r   s       r9   rt   z"SwinTransformerBlock._init_buffersV  sd      ZZ&&--FJJ%%++E**&*FI  iE J	 !r;   r&   rB   rC   c           	      d   t        | j                        r|7|j                  d   |j                  d   }}|j                  }|j                  }n| j
                  \  }}|}|}t        j                  || j                  d   z        | j                  d   z  }t        j                  || j                  d   z        | j                  d   z  }t        j                  d||df||      }d}d| j                  d    f| j                  d    | j                  d    f| j                  d    d ffD ]l  }d| j                  d    f| j                  d    | j                  d    f| j                  d    d ffD ]$  }	||d d |d   |d   |	d   |	d   d d f<   |dz  }& n t        || j                        }
|
j                  d| j                        }
|
j                  d      |
j                  d      z
  }|j                  |dk7  t!        d            j                  |dk(  t!        d            }|S d }|S )Nr   r+   r   )rC   rB   r.   g      Yr   )anyr   r0   rB   rC   r   mathceilr'   rD   zerosr:   r1   r_   r   masked_fillr   )rm   r&   rB   rC   r5   r6   img_maskcnthwmask_windowsr   s               r9   r   z"SwinTransformerBlock.get_attn_mask^  s\    t}wwqz1771:1,,1		!d..q112T5E5Ea5HHA		!d..q112T5E5Ea5HHA{{Aq!Q<uVLHC))!,,-&&q))DOOA,>+>?ooa(($/  T--a001**1--0B/BC//!,,d3 A
 <?HQ!QqT	1Q4!9a781HC ,Hd6F6FGL',,R1A1ABL$..q1L4J4J14MMI!--i1neFmLXXYbfgYginoristI  Ir;   target_window_sizer   c                    t        |      }|(| j                  }t        |      r|d   dz  |d   dz  f}nt        |      }| j                  r||fS t	        | j
                  |      D cg c]  \  }}||k  r|n| }}}t	        | j
                  ||      D cg c]  \  }}}||k  rdn| }}}}t        |      t        |      fS c c}}w c c}}}w )Nr   r+   r   )r   r   r   r   zipr   tuple)rm   r   r   rr   r'   sr   s           r9   r   z'SwinTransformerBlock._calc_window_shift  s    
 ''9:$ $ 6 6$%%7%:a%?ASTUAVZ[A[$\! )*; <  %'88869$:O:OQc6dedaAFq)ee8;D<Q<QS^`q8rssWQ116aq(s
s[!5#444 fss   *C	C	feat_sizec                    || _         ||| _        | j                  |      \  | _        | _        | j                  d   | j                  d   z  | _        | j                  j                  | j                         | j                  | j                  j                  nd}| j                  | j                  j                  nd}| j                  d| j                  rdn| j                  ||      d       y)z
        Args:
            feat_size: New input resolution
            window_size: New window size
            always_partition: Change always_partition attribute if not None
        Nr   r   r   rA   FrZ   )r   r   r   r'   r   r_   r   r   r   rB   rC   re   r   r   )rm   r   r'   r   rB   rC   s         r9   set_input_sizez#SwinTransformerBlock.set_input_size  s     !*'$4D!,0,C,CK,P)$/++A.1A1A!1DD		!!$"2"23*...*D&&$(,(B$$%%D4+=+=VSX+=+Y 	 	
r;   c           	         |j                   \  }}}}t        | j                        }|r7t        j                  || j                  d    | j                  d    fd      }n|}| j
                  d   || j
                  d   z  z
  | j
                  d   z  }| j
                  d   || j
                  d   z  z
  | j
                  d   z  }	t        j                  j                  j                  |ddd|	d|f      }|j                   \  }
}}}
t        || j
                        }|j                  d| j                  |      }t        | dd      r| j                  |      }n| j                  }| j                  ||      }|j                  d| j
                  d   | j
                  d   |      }t!        || j
                  ||      }|d d d |d |d d f   j#                         }|r$t        j                  || j                  d      }|S |}|S )	Nr   r   )r   r+   )shiftsdimsr.   r   F)r   )r0   r   r   rD   rollr'   ra   r   padr:   r1   r_   getattrr   r   r   r=   r3   )rm   r&   r4   r5   r6   r7   	has_shift	shifted_xpad_hpad_w_HpWp	x_windowsr   attn_windowss                   r9   _attnzSwinTransformerBlock._attn  s    WW
1a (	

1tq/A.ADOOTUDVCV-W^deII !!!$q4+;+;A+>'>>$BRBRSTBUU!!!$q4+;+;A+>'>>$BRBRSTBUUHH''++I1a57QR	 2r1 %Y0@0@A	NN2t'7'7;	 4/**95IIyyy; $((T-=-=a-@$BRBRSTBUWXY"<1A1A2rJ	a!RaRl+668	 

9T__6JA  Ar;   c                 >   |j                   \  }}}}|| j                  | j                  | j                  |                  z   }|j	                  |d|      }|| j                  | j                  | j                  |                  z   }|j	                  ||||      }|S )zForward pass.

        Args:
            x: Input features with shape (B, H, W, C).

        Returns:
            Output features with shape (B, H, W, C).
        r.   )r0   r   r   r   r   r   r   r   )rm   r&   r4   r5   r6   r7   s         r9   r   zSwinTransformerBlock.forward  s     WW
1a

4::a= 9::IIaQA 788IIaAq!r;   c                 $    | j                          yr   r   ru   s    r9   r   z0SwinTransformerBlock.init_non_persistent_buffers  r   r;   r   )NNNr   )r   r   r   r   ra   GELU	LayerNormr   r   r   r   r   r   Moduler^   rl   rt   rD   r   rB   rC   r   r
   r	   r   r   r   r   r   r   r   s   @r9   r   r      s.    &*-.%*!&!!!!!)+*,,,%K K  0K  	K 
 smK  +K  K  #K  K  K  K  K  K  K  BIIK   RYY!K ZK )--1+/	&%& U\\*& EKK(	&
 
%,,	&V HL5 %c5c?&: ;5  (c5c?.B(CD5 
uS#Xc3h/	0	52 04	
S#X
 sCx
 'tn	
4%N %,,  r;   r   c                        e Zd ZdZdej
                  ddfdedee   deej                     f fdZ
dej                  dej                  fd	Z xZS )
PatchMergingzVPatch Merging Layer.

    Downsample features by merging 2x2 neighboring patches.
    NrQ   out_dimr   c                     ||d}t         |           || _        |xs d|z  | _         |d|z  fi || _        t        j                  d|z  | j                  fddi|| _        y)z
        Args:
            dim: Number of input channels.
            out_dim: Number of output channels (or 2 * dim if None)
            norm_layer: Normalization layer.
        rA   r+   r,   r\   FN)r]   r^   rQ   r   normra   rf   	reduction)rm   rQ   r   r   rB   rC   rn   rp   s          r9   r^   zPatchMerging.__init__  sj     /)!c'q3w-"-	1s7DLLKuKKr;   r&   r(   c                 h   |j                   \  }}}}ddd|dz  d|dz  f}t        j                  j                  ||      }|j                   \  }}}}|j	                  ||dz  d|dz  d|      j                  dddddd      j                  d      }| j                  |      }| j                  |      }|S )zForward pass.

        Args:
            x: Input features with shape (B, H, W, C).

        Returns:
            Output features with shape (B, H//2, W//2, out_dim).
        r   r+   r   r*   r,   r-   )	r0   ra   r   r   r   r2   rH   r   r   )rm   r&   r4   r5   r6   r7   
pad_valuesr   s           r9   r   zPatchMerging.forward  s     WW
1aAq1uaQ/
MMa,WW
1aIIaaAFAq199!Q1aKSSTUVIIaLNN1r;   )r   r   r   r   ra   r   r   r   r   r   r^   rD   r   r   r   r   s   @r9   r   r     sf     &**,,,LL c]L RYY	L* %,, r;   r   c            "       0    e Zd ZdZdddddddddddej
                  ddfd	ed
edeeef   dededede	e   de
dededededededeee   ef   deej                     f  fdZ	 ddeeef   dede	e   fdZdej&                  dej&                  fdZ xZS ) SwinTransformerStagez|A basic Swin Transformer layer for one stage.

    Contains multiple Swin Transformer blocks and optional downsampling.
    Tr,   Nr   Fr   r   rQ   r   r   depth
downsamplerR   rS   r'   r   r   r   rT   rV   rU   r   r   c                 H   ||d}t         |           || _        || _        |rt	        d |D              n|| _        || _        d| _        t        |      }t	        |D cg c]  }|dz  	 c}      }|rt        d	|||d|| _
        n ||k(  sJ t        j                         | _
        t        j                  t        |      D cg c]E  }t        d	|| j
                  ||||dz  dk(  rdn||	|
||||t!        |t"              r||   n||d|G c} | _        yc c}w c c}w )
a  
        Args:
            dim: Number of input channels.
            out_dim: Number of output 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.
        rA   c              3   &   K   | ]	  }|d z    ywr+   Nr   .0is     r9   	<genexpr>z0SwinTransformerStage.__init__.<locals>.<genexpr>Q  s     &H!qAv&H   Fr+   )rQ   r   r   r   )rQ   r   rR   rS   r'   r   r   r   r   rT   rV   rU   r   r   Nr   )r]   r^   rQ   r   r   output_resolutionr   grad_checkpointingr   r   r   ra   r   
Sequentialranger   
isinstancelistblocks)rm   rQ   r   r   r   r   rR   rS   r'   r   r   r   rT   rV   rU   r   r   rB   rC   rn   r   r   r  rp   s                          r9   r^   zSwinTransformerStage.__init__'  sR   L / 0LV&H7G&H!H\l
"',K8qAF89
 * % 	DO '>!> kkmDO mm$ 5\%&#$ # ! !%!7!7#!'!"Q!1*!1)#!##*4Y*E)A,9% &# $ 9&#s    DA
Dr   c                     || _         t        | j                  t        j                        r|| _        nt        d |D              | _        | j                  D ]   }|j                  | j
                  ||       " y)a   Updates the resolution, window size and so the pair-wise relative positions.

        Args:
            feat_size: New input (feature) resolution
            window_size: New window size
            always_partition: Always partition / shift the window
        c              3   &   K   | ]	  }|d z    ywr  r   r  s     r9   r  z6SwinTransformerStage.set_input_size.<locals>.<genexpr>  s     *Ea16*Er  r   r'   r   N)	r   r  r   ra   r   r	  r   r  r   )rm   r   r'   r   blocks        r9   r   z#SwinTransformerStage.set_input_sizex  sn     !*door{{3%.D"%**E9*E%ED"[[ 	E  00'!1 ! 	r;   r&   r(   c                     | j                  |      }| j                  r6t        j                  j	                         st        | j                  |      }|S | j                  |      }|S )zsForward pass.

        Args:
            x: Input features.

        Returns:
            Output features.
        )r   r
  rD   r   is_scriptingr   r  rm   r&   s     r9   r   zSwinTransformerStage.forward  sU     OOA""599+A+A+Ct{{A.A  AAr;   r   )r   r   r   r   ra   r   r   r	   r   r   r   r   r
   r   r   r   r^   r   rD   r   r   r   r   s   @r9   r   r   !  s\     $&*-.%*!&!!!!35*,,,'O$O$ O$ $CHo	O$
 O$ O$ O$ smO$ +O$ #O$ O$ O$ O$ O$ O$  T%[%/0!O$" RYY#O$j 04	S#X  'tn	2 %,, r;   r   c            ,           e Zd ZdZdddddddd	d
dddddddddeej                  dd
d
fdedededede	dede
edf   de
edf   dee   dedededed ed!ed"ed#ed$ed%eej                     d&ee	eej                     f   d'e	f* fd(Zej&                  j(                  dDd)e	d*ed+d
fd,       Zej&                  j(                  d+ee	   fd-       Z	 	 	 	 	 dEdee
eef      dee
eef      dee
eef      d.edee   d+d
fd/Zej&                  j(                  dFd0ed+ee	ef   fd1       Zej&                  j(                  dGd2ed+d
fd3       Zej&                  j(                  d+ej                  fd4       ZdHdedee	   d+d
fd5Z	 	 	 	 	 dId6ej>                  d7eeee e   f      d8ed9ed:e	d;ed+ee ej>                     e
ej>                  e ej>                     f   f   fd<Z!	 	 	 dJd7eee e   f   d=ed>ed+e e   fd?Z"d6ej>                  d+ej>                  fd@Z#dFd6ej>                  dAed+ej>                  fdBZ$d6ej>                  d+ej>                  fdCZ% xZ&S )Kr%   zSwin 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   FTr   r   g? img_size
patch_sizein_chansnum_classesglobal_pool	embed_dimdepths.rR   rS   r'   r   strict_img_sizer   rT   	drop_rateproj_drop_rateattn_drop_ratedrop_path_rateembed_layerr   weight_initc                    t         !|           ||d}|dv sJ || _        || _        || _        d| _        t        |      | _        || _        t        |d| j                  dz
  z  z        x| _
        | _        g | _        t        |t        t        f      s1t!        | j                        D cg c]  }t        |d|z  z         }} |d"||||d   ||dd|| _        | j"                  j$                  } t'        | j                        |	      }	t        |
t        t        f      s t'        | j                        |
      }
nt        |
      dk(  r|
f| j                  z  }
t        |
      | j                  k(  sJ  t'        | j                        |      }t)        ||d	      }g }|d   }d}t!        | j                        D ]  }||   } |t+        d"i d
|d| d|d   |z  |d   |z  fd||   d|dkD  d||   d|	|   d|
|   d|d| d||   d|d|d|d||   d||gz  }| }|dkD  r|dz  }| xj                  t-        | ||z  d|       gz  c_         t/        j0                  | | _         || j                  fi || _        t7        | j                  |f||| j
                  d|| _        |dk(  rdn|| _        |dk7  r| j=                  d        y!y!c c}w )#a~  
        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.
            embed_layer: Patch embedding layer.
            norm_layer (nn.Module): Normalization layer.
        rA   )r!  r  NHWCr+   r   r   )r"  r#  r$  r'  r   r)  
output_fmtT)	stagewiserQ   r   r   r   r   rR   rS   r'   r   r   r   rT   rV   rU   r   r   layers.)num_chsr   module)	pool_typer*  	input_fmtskipresetF)needs_resetNr   )r]   r^   r%  r$  r&  r2  len
num_layersr'  r   num_featureshead_hidden_sizefeature_infor  r   r  r  patch_embed	grid_sizer   r   r   dictra   r  layersr   r   headweight_init_modeinit_weights)"rm   r"  r#  r$  r%  r&  r'  r(  rR   rS   r'   r   r)  r   rT   r*  r+  r,  r-  r.  r   r/  rB   rC   kwargsrn   r  
patch_griddprrD  in_dimr`   r   rp   s"                                    r9   r^   zSwinTransformer.__init__  s   \ 	/k)))& & f+"47	A$//\]J]D^8^4__D1)eT]3:?:PQQYa/0QIQ ' 	
!l!+	
 	
 %%//
 .9T__-h7+e}54)DOO4[AK"&.4??:K;4??222.Idoo.y9	'$O1t' 	mAlG+  qMU*qMU*" Qi q5 $A, "! (N "2 "10 $A, "  )!" )#$ a&%& &)  F, F1u
$w*uBT_fghfi]j"k!ll7	m8 mmV,t007B7	"
 "oo
 
	 ,7&+@k& %0 !K Rs   ,K!moder;  r(   c                     |xs | j                   }|dv sJ d|v r t        j                  | j                         nd}t	        t        |||      |        y)a*  Initialize model weights.

        Args:
            mode: Weight initialization mode ('jax', 'jax_nlhb', 'moco', or '').
            needs_reset: If True, call reset_parameters() on modules that have it.
                Set to False when modules have already self-initialized in __init__.
        )jaxjax_nlhbmocor:  r!  nlhbr   )	head_biasr;  N)rF  r   logr%  r    r$   )rm   rL  r;  rR  s       r9   rG  zSwinTransformer.init_weights.  sX     ,t,,????39T>TXXd..//r	(P[\^bcr;   c                 v    t               }| j                         D ]  \  }}d|v s|j                  |        |S )z,Parameters that should not use weight decay.rd   )setnamed_parametersadd)rm   nwdnr   s       r9   no_weight_decayzSwinTransformer.no_weight_decay<  sA     e))+ 	DAq-2
	 
r;   window_ratioc                 Z   ||3| j                   j                  ||       | j                   j                  }|t        D cg c]  }||z  	 c}      }t	        | j
                        D ]9  \  }}	dt        |dz
  d      z  }
|	j                  d   |
z  |d   |
z  f||       ; yc c}w )a  Update the image resolution and window size.

        Args:
            img_size: New input resolution, if None current resolution is used.
            patch_size: New patch size, if None use current patch size.
            window_size: New window size, if None based on new_img_size // window_div.
            window_ratio: Divisor for calculating window size from grid size.
            always_partition: Always partition into windows and shift (even if window size < feat size).
        N)r"  r#  r+   r   r   r  )rA  r   rB  r   	enumeraterD  max)rm   r"  r#  r'   r[  r   rI  pgindexstagestage_scales              r9   r   zSwinTransformer.set_input_sizeE  s    " :#9++X*+U))33Jj I|!3 IJK%dkk2 	LE5s519a00K  %a=K7A+9UV'!1 ! 	 !Js   B(coarsec                 2    t        d|rd      S g d      S )z"Group parameters for optimization.z^patch_embedz^layers\.(\d+)))z^layers\.(\d+).downsample)r   )z^layers\.(\d+)\.\w+\.(\d+)N)z^norm)i )stemr  )rC  )rm   rc  s     r9   group_matcherzSwinTransformer.group_matchere  s)      (.$
 	
5
 	
r;   enablec                 4    | j                   D ]	  }||_         y)z)Enable or disable gradient checkpointing.N)rD  r
  )rm   rg  ls      r9   set_grad_checkpointingz&SwinTransformer.set_grad_checkpointingq  s      	*A#)A 	*r;   c                 .    | j                   j                  S )zGet the classifier head.)rE  fcru   s    r9   get_classifierzSwinTransformer.get_classifierw  s     yy||r;   c                 L    || _         | j                  j                  ||       y)zReset the classifier head.

        Args:
            num_classes: Number of classes for new classifier.
            global_pool: Global pooling type.
        )r7  N)r%  rE  r:  )rm   r%  r&  s      r9   reset_classifierz SwinTransformer.reset_classifier|  s      '		{;r;   r&   indicesr   
stop_earlyr2  intermediates_onlyc                 >   |dv sJ d       g }t        t        | j                        |      \  }}	| j                  |      }t        | j                        }
t        j
                  j                         s|s| j                  }n| j                  d|	dz    }t        |      D ]c  \  }} ||      }||v s|r||
dz
  k(  r| j                  |      }n|}|j                  dddd      j                         }|j                  |       e |r|S | j                  |      }||fS )aK  Forward features that returns intermediates.

        Args:
            x: Input image tensor.
            indices: Take last n blocks if int, all if None, select matching indices if sequence.
            norm: Apply norm layer to compatible intermediates.
            stop_early: Stop iterating over blocks when last desired intermediate hit.
            output_fmt: Shape of intermediate feature outputs.
            intermediates_only: Only return intermediate features.

        Returns:
            List of intermediate features or tuple of (final features, intermediates).
        )NCHWzOutput shape must be NCHW.Nr   r   r*   r+   )r   r<  rD  rA  rD   r   r  r]  r   r2   r3   append)rm   r&   rp  r   rq  r2  rr  intermediatestake_indices	max_index
num_stagesstagesr  ra  x_inters                  r9   forward_intermediatesz%SwinTransformer.forward_intermediates  s   , Y&D(DD&"6s4;;7G"Qi Q%
99!!#:[[F[[)a-0F!&) 	.HAuaAL Aa/"iilGG!//!Q15@@B$$W-	.   IIaL-r;   
prune_norm
prune_headc                     t        t        | j                        |      \  }}| j                  d|dz    | _        |rt        j                         | _        |r| j                  dd       |S )aE  Prune layers not required for specified intermediates.

        Args:
            indices: Indices of intermediate layers to keep.
            prune_norm: Whether to prune normalization layer.
            prune_head: Whether to prune the classifier head.

        Returns:
            List of indices that were kept.
        Nr   r   r!  )r   r<  rD  ra   r   r   ro  )rm   rp  r}  r~  rw  rx  s         r9   prune_intermediate_layersz)SwinTransformer.prune_intermediate_layers  s]      #7s4;;7G"Qikk.9q=1DI!!!R(r;   c                 l    | j                  |      }| j                  |      }| j                  |      }|S )z/Forward pass through feature extraction layers.)rA  rD  r   r  s     r9   forward_featuresz SwinTransformer.forward_features  s1    QKKNIIaLr;   
pre_logitsc                 N    |r| j                  |d      S | j                  |      S )zForward pass through classifier head.

        Args:
            x: Feature tensor.
            pre_logits: Return features before final classifier.

        Returns:
            Output tensor.
        T)r  )rE  )rm   r&   r  s      r9   forward_headzSwinTransformer.forward_head  s&     1;tyyty,L		!Lr;   c                 J    | j                  |      }| j                  |      }|S )zoForward pass.

        Args:
            x: Input tensor.

        Returns:
            Output logits.
        )r  r  r  s     r9   r   zSwinTransformer.forward  s)     !!!$a r;   )r!  T)NNN   NF)Tr   )NFFrt  F)r   FT)'r   r   r   r   r   ra   r   r   r   strr	   r   r   r   r   r   r
   r^   rD   r   ignorerG  r   rZ  r   r   r   rf  rj  rm  ro  r   r   r|  r  r  r  r   r   r   s   @r9   r%   r%     s0    +.#$&2)7&*-.%*$(!!!$&$&$'+568ll!1B1'B1 B1 	B1
 B1 B1 B1 #s(OB1 S#XB1 smB1 +B1 #B1 "B1 B1 B1  !B1" "#B1$ "%B1& "'B1( bii)B1* c4		?23+B1, -B1H YYd d d d d YYS   374859 !/3uS#X/ !sCx1 "%S/2	
  'tn 
@ YY	
D 	
T#s(^ 	
 	
 YY*T *T * *
 YY		  <C <hsm <W[ < 8<$$',1 ||1  eCcN341  	1 
 1  1  !%1  
tELL!5tELL7I)I#JJ	K1 j ./$#	3S	>*  	
 
c0%,, 5<< 
Mell 
M 
M 
M %,, r;   
state_dictmodelc                 2   d}d| v rd}ddl }i }| j                  d|       } | j                  d|       } | j                         D ]K  \  }}t        dD cg c]  }||v  c}      r#d	|v re|j                  j
                  j                  j                  \  }}}	}
|j                  d
   |	k7  s|j                  d   |
k7  rt        ||	|
fddd      }|j                  d      r|j                  |dd       }|j                  |j                  j                  k7  s|j                  d   |j                  d   k7  r,t        ||j                  |j                  j                        }|r&|j                  dd |      }|j                  dd      }|||<   N |S c c}w )zConvert patch embedding weight from manual patchify + linear proj to conv.

    Args:
        state_dict: State dictionary from checkpoint.
        model: Model instance.

    Returns:
        Filtered state dictionary.
    Tzhead.fc.weightFr   Nr  r  )rY   r   zpatch_embed.proj.weightr   r.   bicubic)interpolation	antialiasverboserd   ir   r{   zlayers.(\d+).downsamplec                 D    dt        | j                  d            dz    dS )Nr4  r   z.downsample)r   group)r&   s    r9   <lambda>z&checkpoint_filter_fn.<locals>.<lambda>  s$    ws177ST:YZGZF[[f=g r;   zhead.zhead.fc.)regetitemsr   rA  ri   r   r0   r   endswithget_submodulerd   r'   r   subreplace)r  r  old_weightsr  out_dictr   r   rY  r   r5   r6   ms               r9   checkpoint_filter_fnr    s    K:%H4Jj9J  " 1 HI1QIJ$)**//66<<JAq!Qwwr{a1772;!#3(F"+"  ::45##AdsG,Aww!88>>>!--PQBRVWVcVcdeVfBf-$%MM#$#A#A#G#G 13gijkA		':.A9: O9 Js   F
variant
pretrainedc           	          t        d t        |j                  dd            D              }|j                  d|      }t	        t
        | |ft        t        d|      d|}|S )zCreate a Swin Transformer model.

    Args:
        variant: Model variant name.
        pretrained: Load pretrained weights.
        **kwargs: Additional model arguments.

    Returns:
        SwinTransformer model instance.
    c              3   &   K   | ]	  \  }}|  y wr   r   )r  r  r   s      r9   r  z+_create_swin_transformer.<locals>.<genexpr>.  s     \da\r  r(  )r   r   r*   r   out_indicesT)flatten_sequentialr  )pretrained_filter_fnfeature_cfg)r   r]  r  popr   r%   r  rC  )r  r  rH  default_out_indicesr  r  s         r9   _create_swin_transformerr  #  sj      \i

8\8Z.[\\**],?@K *1DkJ 	E Lr;   urlc                 4    | ddddddt         t        ddd	d
|S )z9Create default configuration for Swin Transformer models.r  )r*   r  r  )r   r   g?r  Tzpatch_embed.projzhead.fcmit)r  r%  
input_size	pool_sizecrop_pctr  fixed_input_sizemeanrs   
first_conv
classifierlicenser   )r  rH  s     r9   _cfgr  :  s7     =v%.B(	 # r;   z.swin_small_patch4_window7_224.ms_in22k_ft_in1kztimm/zvhttps://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_small_patch4_window7_224_22kto1k_finetune.pth)	hf_hub_idr  z-swin_base_patch4_window7_224.ms_in22k_ft_in1kzlhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window7_224_22kto1k.pthz.swin_base_patch4_window12_384.ms_in22k_ft_in1kzmhttps://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  z.swin_large_patch4_window7_224.ms_in22k_ft_in1kzmhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window7_224_22kto1k.pthz/swin_large_patch4_window12_384.ms_in22k_ft_in1kznhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window12_384_22kto1k.pthz$swin_tiny_patch4_window7_224.ms_in1kzdhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_tiny_patch4_window7_224.pthz%swin_small_patch4_window7_224.ms_in1kzehttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_small_patch4_window7_224.pthz$swin_base_patch4_window7_224.ms_in1kzdhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window7_224.pthz%swin_base_patch4_window12_384.ms_in1kzehttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window12_384.pthz-swin_tiny_patch4_window7_224.ms_in22k_ft_in1kzuhttps://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_tiny_patch4_window7_224_22kto1k_finetune.pthz%swin_tiny_patch4_window7_224.ms_in22kzhhttps://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_tiny_patch4_window7_224_22k.pthiQU  )r  r  r%  z&swin_small_patch4_window7_224.ms_in22kzihttps://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_small_patch4_window7_224_22k.pthz%swin_base_patch4_window7_224.ms_in22kzhhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window7_224_22k.pthz&swin_base_patch4_window12_384.ms_in22kzihttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window12_384_22k.pth)r  r  r  r  r  r%  z&swin_large_patch4_window7_224.ms_in22kzihttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window7_224_22k.pthz'swin_large_patch4_window12_384.ms_in22kzjhttps://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window12_384_22k.pthzswin_s3_tiny_224.ms_in1kzbhttps://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_s3_small_224.ms_in1kzswin_s3_base_224.ms_in1kc           	      R    t        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(  rR   r  )swin_tiny_patch4_window7_224rC  r  r  rH  
model_argss      r9   r  r    sF     R`noJ#&]3=]AEjA[TZA[] ]r;   c           	      R    t        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  r  )swin_small_patch4_window7_224r  r  s      r9   r  r    sF     RaopJ#'^4>^BFzB\U[B\^ ^r;   c           	      R    t        ddddd      }t        	 dd| it        |fi |S )	z Swin-B @ 224x224
    r,   r      r  r,   r         r  r  )swin_base_patch4_window7_224r  r  s      r9   r  r    sF     SbpqJ#&]3=]AEjA[TZA[] ]r;   c           	      R    t        ddddd      }t        	 dd| it        |fi |S )	z Swin-B @ 384x384
    r,   r  r  r  r  r  r  )swin_base_patch4_window12_384r  r  s      r9   r  r    sF     c-cqrJ#'^4>^BFzB\U[B\^ ^r;   c           	      R    t        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  r  )swin_large_patch4_window7_224r  r  s      r9   r  r    sF     SbqrJ#'^4>^BFzB\U[B\^ ^r;   c           	      R    t        ddddd      }t        	 dd| it        |fi |S )	z Swin-L @ 384x384
    r,   r  r  r  r  r  r  )swin_large_patch4_window12_384r  r  s      r9   r  r    sF     c-crsJ#(_5?_CG
C]V\C]_ _r;   c           	      P    t        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  r  )swin_s3_tiny_224r  r  s      r9   r  r    s;     -2l^lnJ#l:lQUV`QkdjQkllr;   c           	      P    t        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  r  )swin_s3_small_224r  r  s      r9   r  r    s;     /RaoqJ#mJmRVWaRlekRlmmr;   c           	      P    t        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  r  )swin_s3_base_224r  r  s      r9   r  r    s;     -2m_moJ#l:lQUV`QkdjQkllr;   )"swin_base_patch4_window7_224_in22k#swin_base_patch4_window12_384_in22k#swin_large_patch4_window7_224_in22k$swin_large_patch4_window12_384_in22kr   r  )r!  )Or   loggingr   typingr   r   r   r   r   r   r	   r
   r   rD   torch.nnra   	timm.datar   r   timm.layersr   r   r   r   r   r   r   r   r   r   r   r   _builderr   	_featuresr   _features_fxr   _manipulater   r    	_registryr!   r"   r#   vision_transformerr$   __all__	getLoggerr   _loggerr   r   r   r:   r=   rM   r   rO   r   r   r   r%   rC  r  r  r   r  r  default_cfgsr  r  r  r  r  r  r  r  r  r   r;   r9   <module>r     s2  
"   O O O   AL L L L * + 3 4 Y Y 4

'

H
%#uS#X./ <<38_ \\& ELL uS#X 3 SV [`[g[g  $#s #3 # #0Tbii Tno299 od-299 -`299 DLbii L^
.T .")) .S%,,EV@W .bc t Ra .	c 	T#s(^ 	 % H&4d E7HH& 4Tz6}	H& 5d{ Hs7DH& 5d{7~H& 6t| Hs8DH&& +Dr-u'H&, ,Ts.v-H&2 +Dr-u3H&8 ,Ts Hs.D9H&D 4T D6FEH&L ,Tv.MH&T -dw/UH&\ ,Tv.]H&d -dw HsPU/WeH&l -dw/mH&t .tx HsPU0WuH&~ p!rH&D "&p"r !%p!rKH& HV ] ] ] ^ ^ ^ ] ] ] ^ ^ ^ ^ ^ ^ _/ _ _ mO m m n_ n n mO m m H*Q+S+S,U	' r;   