
    ^j/                     ~   d Z ddlZddl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ZddlmZ ddlmc 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" ddl#m$Z$ dd	l%m&Z& dd
l'm(Z( ddl)m*Z* ddl+m,Z,m-Z-m.Z. dgZ/ee0ee0e0f   f   Z1dejd                  dee0e0f   dejd                  fdZ3e(dejd                  dee0e0f   dee0e0f   dejd                  fd       Z4 G d dejj                        Z6 G d dejj                        Z7 G d dejj                        Z8 G d dejj                        Z9 G d dejj                        Z:dee;ejd                  f   dejj                  dee;ejd                  f   fd Z<dOd!e;d"e=de:fd#Z>dPd$Z? e, e?d%d&'       e?d%d(d)d*d+,       e?d%d-'       e?d%d.d)d*d+,       e?d%d/'       e?d%d0'       e?d%d1'       e?d%d2'       e?d%d3'       e?d%d4'       e?d%d5d6d7d89       e?d%d:d6d7d89      d;      Z@e-dOd"e=de:fd<       ZAe-dOd"e=de:fd=       ZBe-dOd"e=de:fd>       ZCe-dOd"e=de:fd?       ZDe-dOd"e=de:fd@       ZEe-dOd"e=de:fdA       ZFe-dOd"e=de:fdB       ZGe-dOd"e=de:fdC       ZHe-dOd"e=de:fdD       ZIe-dOd"e=de:fdE       ZJe-dOd"e=de:fdF       ZKe-dOd"e=de:fdG       ZL e.eMdHdIdJdKdLdMdN       y)QaK   Swin Transformer V2
A PyTorch impl of : `Swin Transformer V2: Scaling Up Capacity and Resolution`
    - https://arxiv.org/abs/2111.09883

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

Modifications and additions for timm hacked together by / Copyright 2022, Ross Wightman
    N)partial)	AnyCallableDictListOptionalSetTupleTypeUnionIMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STD)
PatchEmbedMlpDropPathcalculate_drop_path_rates	to_2tupletrunc_normal_ClassifierHeadresample_patch_embedndgridget_act_layer	LayerType   )build_model_with_cfg)feature_take_indices)register_notrace_function)
checkpoint)generate_default_cfgsregister_modelregister_model_deprecationsSwinTransformerV2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 tensor of shape (B, H, W, C).
        window_size: Window size (height, width).

    Returns:
        Windows tensor of shape (num_windows*B, window_size[0], window_size[1], C).
    r   r               shapeviewpermute
contiguous)r$   r%   BHWCwindowss          j/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/swin_transformer_v2.pywindow_partitionr8   %   s     JAq!Q	q!{1~%{1~qKN7JKXYN\]^Aii1aAq)446;;BAP[\]P^`abGN    r6   img_sizec                     |\  }}| j                   d   }| j                  d||d   z  ||d   z  |d   |d   |      }|j                  dddddd      j                         j                  d|||      }|S )a1  Merge windows back to feature map.

    Args:
        windows: Windows tensor of shape (num_windows * B, window_size[0], window_size[1], C).
        window_size: Window size (height, width).
        img_size: Image size (height, width).

    Returns:
        Feature map tensor of shape (B, H, W, C).
    r,   r   r   r(   r)   r*   r+   r-   )r6   r%   r:   r3   r4   r5   r$   s          r7   window_reverser<   8   s      DAqbARk!n,a;q>.A;q>S^_`SacdeA			!Q1a#..055b!QBAHr9   c                   :    e Zd ZdZ	 	 	 	 	 	 	 ddedeeef   dedededed	ed
eeef   ddf fdZddZ	ddZ
	 	 ddeej                  ej                  f   fdZdeeef   ddfdZddZddej                  deej                     dej                  fdZ xZS )WindowAttentionzWindow based multi-head self attention (W-MSA) module with relative position bias.

    Supports both shifted and non-shifted window attention with continuous relative
    position bias and cosine attention.
    Ndimr%   	num_headsqkv_biasqkv_bias_separate	attn_drop	proj_droppretrained_window_sizer&   c           	         |	|
d}t         |           || _        || _        t	        |      | _        || _        || _        t        j                  t        j                  |ddffi |      | _        t        j                  t        j                  dddi|t        j                  d      t        j                  d|fddi|      | _        t        j                  ||d	z  fddi|| _        |rt        j                  t        j                  |fi |      | _        | j'                  d
t        j                  |fi |d       t        j                  t        j                  |fi |      | _        nd| _        d| _        d| _        t        j,                  |      | _        t        j                  ||fi || _        t        j,                  |      | _        t        j4                  d      | _        | j                  \  }}| j'                  dt        j                  dd|z  dz
  d|z  dz
  dfi |d       | j'                  dt        j                  ||z  ||z  |	t        j8                        d       | j;                          y)a4  Initialize window attention module.

        Args:
            dim: Number of input channels.
            window_size: The height and width of the window.
            num_heads: Number of attention heads.
            qkv_bias: If True, add a learnable bias to query, key, value.
            qkv_bias_separate: If True, use separate bias for q, k, v projections.
            attn_drop: Dropout ratio of attention weight.
            proj_drop: Dropout ratio of output.
            pretrained_window_size: The height and width of the window in pre-training.
        devicedtyper   r)      biasT)inplaceFr(   k_bias
persistentNr,   r?   relative_coords_tablerelative_position_index)r)   rJ   )super__init__r?   r%   r   rE   r@   rB   nn	Parametertorchemptylogit_scale
SequentialLinearReLUcpb_mlpqkvq_biasregister_bufferv_biasrM   DropoutrC   projrD   Softmaxsoftmaxlongreset_parameters)selfr?   r%   r@   rA   rB   rC   rD   rE   rH   rI   ddwin_hwin_w	__class__s                 r7   rT   zWindowAttention.__init__V   s.   2 /&&/0F&G#"!2<<Y14E(L(LM }}II.4.2.GGD!IIc9757B7
 99S#'<<<,,u{{3'="'=>DK  5;;s+Ab+Ae T,,u{{3'="'=>DKDKDKDKI.IIc3-"-	I.zzb) ''u#KK1u9q=!e)a-AbA 	 	

 	%KKuu}V5::V 	 	
 	r9   c                 Z   t         j                  j                  | j                  t	        j
                  d             | j                  Rt         j                  j                  | j                         t         j                  j                  | j                         | j                          y)z"Initialize parameters and buffers.
   N)
rU   init	constant_rY   mathlogr_   zeros_ra   _init_buffersrh   s    r7   rg   z WindowAttention.reset_parameters   sb    
$**DHHRL9;;"GGNN4;;'GGNN4;;'r9   c                 `   | j                   | j                   j                          | j                  | j                  j                  j
                  | j                  j                  j                        \  }}| j                  j                  |       | j                  j                  |       y)z.Compute and fill non-persistent buffer values.NrG   )
rM   zero_"_make_pair_wise_relative_positionsrc   weightrH   rI   rQ   copy_rR   )rh   rQ   rR   s      r7   rt   zWindowAttention._init_buffers   s    ;;"KK9=9`9`99##**$))2B2B2H2H :a :
66 	""(()>?$$**+BCr9   c                    t        j                  | j                  d   dz
   | j                  d   |t         j                        }t        j                  | j                  d   dz
   | j                  d   |t         j                        }t        j                  t        ||            }|j                  ddd      j                         j                  d      }| j                  d   dkD  rO|dddddddfxx   | j                  d   dz
  z  cc<   |dddddddfxx   | j                  d   dz
  z  cc<   nN|dddddddfxx   | j                  d   dz
  z  cc<   |dddddddfxx   | j                  d   dz
  z  cc<   |dz  }t        j                  |      t        j                  t        j                  |      dz         z  t        j                  d      z  }|j                  |      }t        j                  | j                  d   |t         j                        }t        j                  | j                  d   |t         j                        }t        j                  t        ||            }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   | j                  d   dz
  z  cc<   |
dddddfxx   | j                  d   dz
  z  cc<   |
dddddfxx   d| j                  d   z  dz
  z  cc<   |
j#                  d	      }||fS )
zCompute pair-wise relative position index and coordinates table.

        Returns:
            Tuple of (relative_coords_table, relative_position_index)
        r   r   rG   r)   N         ?)rI   r,   )rW   aranger%   float32stackr   r0   r1   	unsqueezerE   signlog2absrq   torf   flattensum)rh   rH   rI   relative_coords_hrelative_coords_wrQ   coords_hcoords_wcoordscoords_flattenrelative_coordsrR   s               r7   rx   z2WindowAttention._make_pair_wise_relative_positions   s    "LLq!A%&(8(8(;FRWR_R_a!LLq!A%&(8(8(;FRWR_R_a %F3DFW,X Y 5 = =aA F Q Q S ] ]^_ `&&q)A-!!Q1*-$2M2Ma2PST2TU-!!Q1*-$2M2Ma2PST2TU-!!Q1*-$2B2B12E2IJ-!!Q1*-$2B2B12E2IJ-" %

+@ AEJJII+,s2E4 !46:iil!C 5 8 8u 8 E << 0 0 3F%**U<< 0 0 3F%**UVHh78vq1(At4~aqj7QQ)11!Q:EEG1a D$4$4Q$7!$;; 1a D$4$4Q$7!$;; 1a A(8(8(;$;a$?? "1"5"5b"9$&===r9   c                 8   t        |      }|| j                  k7  r| j                  J | j                  j                  }| j                  j                  }|| _        | j                  ||      \  }}| j                  d|d       | j                  d|d       yy)zUpdate window size and regenerate relative position tables.

        Args:
            window_size: New window size (height, width).
        NrG   rQ   FrN   rR   )r   r%   rQ   rH   rI   rx   r`   )rh   r%   rH   rI   rQ   rR   s         r7   set_window_sizezWindowAttention.set_window_size   s      ,$***--999//66F..44E*D77vU7S ;!#:  !8:O\a b  !:<S`e f +r9   c                 $    | j                          y)z"Initialize non-persistent buffers.N)rt   ru   s    r7   init_non_persistent_buffersz+WindowAttention.init_non_persistent_buffers   s    r9   r$   maskc                    |j                   \  }}}| j                  | j                  |      }nt        j                  | j                  | j
                  | j                  f      }| j                  r| j                  |      }||z  }n,t        j                  || j                  j                  |      }|j                  ||d| j                  d      j                  ddddd      }|j                  d      \  }}	}
t        j                  |d      t        j                  |	d      j!                  d	d      z  }t        j"                  | j$                  t'        j(                  d
            j+                         }||z  }| j-                  | j.                        j1                  d| j                        }|| j2                  j1                  d         j1                  | j4                  d   | j4                  d   z  | j4                  d   | j4                  d   z  d      }|j                  ddd      j7                         }dt        j8                  |      z  }||j;                  d      z   }||j                   d   }|j1                  d|| j                  ||      |j;                  d      j;                  d      z   }|j1                  d| j                  ||      }| j=                  |      }n| j=                  |      }| j?                  |      }||
z  j!                  dd      j                  |||      }| jA                  |      }| jC                  |      }|S )a#  Forward pass of window attention.

        Args:
            x: Input features with shape of (num_windows*B, N, C).
            mask: Attention mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None.

        Returns:
            Output features with shape of (num_windows*B, N, C).
        )ry   rK   r(   r,   r)   r   r   r*   rP   g      Y@)max   )"r.   r_   r^   rW   catrM   ra   rB   Flinearry   reshaper@   r0   unbind	normalize	transposeclamprY   rq   rr   expr]   rQ   r/   rR   r%   r1   sigmoidr   re   rC   rc   rD   )rh   r$   r   B_Nr5   r^   rA   qkvattnrY   relative_position_bias_tablerelative_position_biasnum_wins                   r7   forwardzWindowAttention.forward   s    77Aq;;((1+Cyy$++t{{DKK!HIH%%hhqkxhhqxHkk"aDNNB7??1aAN**Q-1a A2&QB)?)I)I"b)QQkk$"2"28KLPPRk!'+||D4N4N'O'T'TUWY]YgYg'h$!=d>Z>Z>_>_`b>c!d!i!iQ$"2"21"55t7G7G7JTM]M]^_M`7`bd"f!7!?!?1a!H!S!S!U!#emm4J&K!K,66q99jjmG99R$..!Q?$..QRBSB]B]^_B``D99RA6D<<%D<<%D~~d#AX  A&..r1a8IIaLNN1r9   )TF        r   )r   r   NNr&   N)NNN)__name__
__module____qualname____doc__intr
   boolfloatrT   rg   rt   rW   Tensorrx   r   r   r   r   __classcell__rl   s   @r7   r>   r>   O   s    "&+!!6<F F  sCxF  	F 
 F   $F  F  F  %*#s(OF  
F PD (> 
u||U\\)	*	(>Tg5c? gt g"1 1Xell-C 1u|| 1r9   r>   c                       e Zd ZdZddddddddddej
                  dd	d	fd
ededed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f f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d ee   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   dd	fd#Zdej                   dej                   fd$Zdej                   dej                   fd%Z xZS )(SwinTransformerV2BlockzSwin Transformer V2 Block.

    A standard transformer block with window attention and shifted window attention
    for modeling long-range dependencies efficiently.
       r   F      @Tr   geluNr?   input_resolutionr@   r%   
shift_sizealways_partitiondynamic_mask	mlp_ratiorA   rD   rC   	drop_path	act_layer
norm_layerrE   c                 8   ||d}t         |           || _        t        |      | _        || _        t        |      | _        || _        || _        | j                  ||      \  | _
        | _        | j                  d   | j                  d   z  | _        || _        t        |      }t        |ft        | j                        ||	||
t        |      d|| _         ||fi || _        |dkD  rt%        |      nt'        j(                         | _        t-        d|t/        ||z        ||
d|| _         ||fi || _        |dkD  rt%        |      nt'        j(                         | _        | j7                  d| j                  rdn | j8                  di |d	
       y)a  
        Args:
            dim: Number of input channels.
            input_resolution: Input resolution.
            num_heads: Number of attention heads.
            window_size: Window size.
            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.
            pretrained_window_size: Window size in pretraining.
        rG   r   r   )r%   r@   rA   rC   rD   rE   r   )in_featureshidden_featuresr   drop	attn_maskNFrN    )rS   rT   r?   r   r   r@   target_shift_sizer   r   _calc_window_shiftr%   r   window_arear   r   r>   r   norm1r   rU   Identity
drop_path1r   r   mlpnorm2
drop_path2r`   get_attn_mask)rh   r?   r   r@   r%   r   r   r   r   rA   rD   rC   r   r   r   rE   rH   rI   ri   rl   s                      r7   rT   zSwinTransformerV2Block.__init__*  s   J / )*: ;"!*:!6 0(,0,C,CKQ[,\)$/++A.1A1A!1DD"!),	#	
!$"2"23#,-C#D	
 	
	  *r*
1:R(9-R[[] 
i0	

 
  *r*
1:R(9-R[[]%%D+=4+=+=+C+C 	 	
r9   r$   rH   rI   r&   c           	         t        | j                        r|)t        j                  dg| j                  d||      }nJt        j                  d|j
                  d   |j
                  d   df|j                  |j                        }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 )	zGenerate attention mask for shifted window attention.

        Args:
            x: Input tensor for dynamic shape calculation.

        Returns:
            Attention mask or None if no shift.
        Nr   rG   r)   r   r,   g      Yr   )anyr   rW   zerosr   r.   rH   rI   r%   r8   r/   r   r   masked_fillr   )
rh   r$   rH   rI   img_maskcnthwmask_windowsr   s
             r7   r   z$SwinTransformerV2Block.get_attn_masky  s    ty ;;'ED,A,A'E1'Ef\ab ;;1771:qwwqz1'Eahh^_^e^ef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r9   target_window_sizer   c                    t        |      }|(| j                  }t        |      r|d   dz  |d   dz  f}nt        |      }| j                  r||fS t        |      }t        |      }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 )a  Calculate window size and shift size based on input resolution.

        Args:
            target_window_size: Target window size.
            target_shift_size: Target shift size.

        Returns:
            Tuple of (adjusted_window_size, adjusted_shift_size).
        r   r)   r   )r   r   r   r   zipr   tuple)rh   r   r   rr   r%   sr   s           r7   r   z)SwinTransformerV2Block._calc_window_shift  s    ''9:$ $ 6 6$%%7%:a%?ASTUAVZ[A[$\! )*; <  %'888&'9:%&7869$:O:OQc6dedaAFq)ee8;D<Q<QS^`q8rssWQ116aq(s
s[!5#444 fss    C1C%	feat_sizec                    || _         ||| _        | j                  t        |            \  | _        | _        | 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)zSet input size and update window configuration.

        Args:
            feat_size: New feature map size.
            window_size: New window size.
            always_partition: Override always_partition setting.
        Nr   r   r   rG   FrN   )r   r   r   r   r%   r   r   r   r   r   rH   rI   r`   r   r   )rh   r   r%   r   rH   rI   s         r7   set_input_sizez%SwinTransformerV2Block.set_input_size  s     !*'$4D!,0,C,CIkDZ,[)$/++A.1A1A!1DD		!!$"2"23*...*D&&$(,(B$$%%D4+=+=VSX+=+Y 	 	
r9   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
                  ||f      }|d	d	d	|d	|d	d	f   j#                         }|r$t        j                  || j                  d      }|S |}|S )
zApply windowed attention with optional shift.

        Args:
            x: Input tensor of shape (B, H, W, C).

        Returns:
            Output tensor of shape (B, H, W, C).
        r   r   )r   r)   )shiftsdimsr,   r   F)r   N)r.   r   r   rW   rollr%   rU   
functionalpadr8   r/   r   getattrr   r   r   r<   r1   )rh   r$   r2   r3   r4   r5   	has_shift	shifted_xpad_hpad_w_HpWp	x_windowsr   attn_windowss                   r7   _attnzSwinTransformerV2Block._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"<1A1AB8L	a!RaRl+668	 

9T__6JA  Ar9   c                 >   |j                   \  }}}}|| j                  | j                  | j                  |                  z   }|j	                  |d|      }|| j                  | j                  | j                  |                  z   }|j	                  ||||      }|S )Nr,   )r.   r   r   r   r   r   r   r   )rh   r$   r2   r3   r4   r5   s         r7   r   zSwinTransformerV2Block.forward  s    WW
1a

4::a= 9::IIaQ

488A; 788IIaAq!r9   )NNNr   )r   r   r   r   rU   	LayerNormr   _int_or_tuple_2_tr   r   r   r   ModulerT   r   rW   r   rH   rI   r   r
   r   r   r   r   r   r   s   @r7   r   r   #  s    ./,-%*!&!!!!!#)*,,,89%M
M
 0M
 	M

 +M
 *M
 #M
 M
 M
 M
 M
 M
 M
 !M
 RYYM
  %6!M
b )--1+/	'%' U\\*' EKK(	'
 
%,,	'X >B5 15  ((9:5 
uS#Xc3h/	0	5J 04	
S#X
 sCx
 'tn	

 

8,u|| , ,\ %,, r9   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 )
PatchMergingzPatch Merging Layer.

    Merges 2x2 neighboring patches and projects to higher dimension,
    effectively downsampling the feature maps.
    Nr?   out_dimr   c                     ||d}t         |           || _        |xs d|z  | _        t	        j
                  d|z  | j                  fddi|| _         || j                  fi || _        y)z
        Args:
            dim (int): Number of input channels.
            out_dim (int): Number of output channels (or 2 * dim if None)
            norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm
        rG   r)   r*   rK   FN)rS   rT   r?   r   rU   r[   	reductionnorm)rh   r?   r   r   rH   rI   ri   rl   s          r7   rT   zPatchMerging.__init__  sj     /)!c'1s7DLLKuKKt||2r2	r9   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 )Nr   r)   r   r(   r*   r+   )	r.   rU   r   r   r   r0   r   r  r  )rh   r$   r2   r3   r4   r5   
pad_valuesr   s           r7   r   zPatchMerging.forward2  s    WW
1aAq1uaQ/
MMa,WW
1aIIaaAFAq199!Q1aKSSTUVNN1IIaLr9   )r   r   r   r   rU   r   r   r   r   r   rT   rW   r   r   r   r   s   @r7   r   r     sb     &**,,,33 c]3 RYY	3*
 
%,, 
r9   r   c            '       L    e Zd ZdZdddddddddej
                  ddddfd	ed
edededededededede	dede	de	de	de
eeej                     f   deej                     dededdf& fdZ	 d"deeef   dedee   ddfdZdej&                  dej&                  fd Zd#d!Z xZS )$SwinTransformerV2StagezA Swin Transformer V2 Stage.

    A single stage consisting of multiple Swin Transformer blocks with
    optional downsampling at the beginning.
    Fr   Tr   r   r   Nr?   r   r   depthr@   r%   r   r   
downsampler   rA   rD   rC   r   r   r   rE   output_nchwr&   c                 ^   ||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]F  }t!        d	|| j
                  |||dz  dk(  rdn||||
|||t#        |t$              r||   n||||d|H 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.
            num_heads: Number of attention heads.
            window_size: Local window size.
            always_partition: Always partition into full windows and shift
            dynamic_mask: Create attention mask in forward based on current input size
            downsample: Use downsample layer at start of the block.
            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.
            act_layer: Activation layer type.
            norm_layer: Normalization layer.
            pretrained_window_size: Local window size in pretraining.
            output_nchw: Output tensors on NCHW format instead of NHWC.
        rG   c              3   &   K   | ]	  }|d z    ywr)   Nr   .0is     r7   	<genexpr>z2SwinTransformerV2Stage.__init__.<locals>.<genexpr>v  s     &H!qAv&H   Fr)   )r?   r   r   r   )r?   r   r@   r%   r   r   r   r   rA   rD   rC   r   r   r   rE   Nr   )rS   rT   r?   r   r   output_resolutionr  r	  grad_checkpointingr   r   r  rU   r   
ModuleListranger   
isinstancelistblocks)rh   r?   r   r   r  r@   r%   r   r   r  r   rA   rD   rC   r   r   r   rE   r	  rH   rI   ri   r   r   r  rl   s                            r7   rT   zSwinTransformerV2Stage.__init__F  sQ   X / 0LV&H7G&H!H\l
&"',K8qAF89
 *asGPZa^`aDO'>!> kkmDO mm& 5\'%#& % # !%!7!7#'!"Q!1*!1)#!##*4Y*E)A,9#%'=  !%# $ 9%#s   'D%AD*r   c                 .   || _         t        | j                  t        j                        r|| _        n3t        | j                  t              sJ t        d |D              | _        | j                  D ]   }|j                  | j
                  ||       " y)zUpdate resolution, window size and 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     r7   r  z8SwinTransformerV2Stage.set_input_size.<locals>.<genexpr>  s     *Ea16*Er  r   r%   r   N)
r   r  r  rU   r   r  r   r   r  r   )rh   r   r%   r   blocks        r7   r   z%SwinTransformerV2Stage.set_input_size  s     !*door{{3%.D"doo|<<<%**E9*E%ED"[[ 	E  00'!1 ! 	r9   r$   c                     | j                  |      }| j                  D ]A  }| j                  r+t        j                  j                         st        ||      }: ||      }C |S )zForward pass through the stage.

        Args:
            x: Input tensor of shape (B, H, W, C).

        Returns:
            Output tensor of shape (B, H', W', C').
        )r  r  r  rW   jitis_scriptingr   )rh   r$   blks      r7   r   zSwinTransformerV2Stage.forward  sY     OOA;; 	C&&uyy/E/E/GsA&F		
 r9   c                    | j                   D ]  }t        j                  j                  |j                  j
                  d       t        j                  j                  |j                  j                  d       t        j                  j                  |j                  j
                  d       t        j                  j                  |j                  j                  d        y)z/Initialize residual post-normalization weights.r   N)r  rU   ro   rp   r   rK   ry   r   )rh   r   s     r7   _init_respostnormz(SwinTransformerV2Stage._init_respostnorm  s    ;; 	3CGGciinna0GGcii..2GGciinna0GGcii..2		3r9   r   r   )r   r   r   r   rU   r   r   r   r   r   r   strr   r   rT   r
   r   r   rW   r   r   r"  r   r   s   @r7   r  r  ?  s    &+!&$!!!!!5;*,,,89 %+R$R$ R$ 0	R$
 R$ R$ +R$ #R$ R$ R$ R$ R$ R$ R$ R$  S$ryy/12!R$" RYY#R$$ %6%R$& 'R$, 
-R$p 04	S#X  'tn	
 
4 %,, $3r9   r  c            +       l    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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de
de
ded e
d!ed"ed#ed$ed%eeef   d&eej                     d'e	edf   f( fd(ZdEd)e
d*dfd+ZdEd,ej                  d)e
d*dfd-Z	 	 	 	 	 dFdee	eef      dee	eef      dee	eef      d.ee   dee
   f
d/Zej,                  j.                  d*ee   fd0       Zej,                  j.                  dGd1e
d*eeef   fd2       Zej,                  j.                  dEd3e
d*dfd4       Zej,                  j.                  d*ej                  fd5       ZdHdedee   d*dfd6Z	 	 	 	 	 dId7ej@                  d8eeee!e   f      d9e
d:e
d;ed<e
d*ee!ej@                     e	ej@                  e!ej@                     f   f   fd=Z"	 	 	 dJd8eee!e   f   d>e
d?e
fd@Z#d7ej@                  d*ej@                  fdAZ$dGd7ej@                  dBe
d*ej@                  fdCZ%d7ej@                  d*ej@                  fdDZ& xZ'S )Kr#   a   Swin Transformer V2.

    A hierarchical vision transformer using shifted windows for efficient
    self-attention computation with continuous position bias.

    A PyTorch impl of : `Swin Transformer V2: Scaling Up Capacity and Resolution`
        - https://arxiv.org/abs/2111.09883
       r*   r(     avg`   r)   r)      r)   r(   r*        r   FTr   r   g?r   )r   r   r   r   Nr:   
patch_sizein_chansnum_classesglobal_pool	embed_dimdepths.r@   r%   r   strict_img_sizer   rA   	drop_rateproj_drop_rateattn_drop_ratedrop_path_rater   r   pretrained_window_sizesc                 d   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         }}t#        d"||||d   ||dd|| _        | j$                  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|d||   |gz  }|}|dkD  r|dz  }| xj                  t-        |d|z  d|       gz  c_         t/        j0                  | | _         || j                  fi || _        t7        | j                  |f||| j
                  d|| _        | j;                  d        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 stage (layer).
            num_heads: Number of attention heads in different layers.
            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: Head dropout rate.
            proj_drop_rate: Projection dropout rate.
            attn_drop_rate: Attention dropout rate.
            drop_path_rate: Stochastic depth rate.
            norm_layer: Normalization layer.
            act_layer: Activation layer type.
            patch_norm: If True, add normalization after patch embedding.
            pretrained_window_sizes: Pretrained window sizes of each layer.
            output_fmt: Output tensor format if not None, otherwise output 'NHWC' by default.
        rG   ) r'  NHWCr)   r   r   )r:   r.  r/  r2  r   r4  
output_fmtT)	stagewiser?   r   r   r  r  r@   r%   r   r   r   rA   rD   rC   r   r   r   rE   r*   layers.)num_chsr  module)	pool_typer5  	input_fmtFneeds_resetNr   )rS   rT   r0  r/  r1  r=  len
num_layersr2  r   num_featureshead_hidden_sizefeature_infor  r   r  r  r   patch_embed	grid_sizer   r  dictrU   rZ   layersr  r   headinit_weights)!rh   r:   r.  r/  r0  r1  r2  r3  r@   r%   r   r4  r   rA   r5  r6  r7  r8  r   r   r9  rH   rI   kwargsri   r  rL  dprrN  in_dimscaler   rl   s!                                   r7   rT   zSwinTransformerV2.__init__  s   ` 	/& k)))& f+"47	A$//\]J]D^8^4__D1)eT]3:?:PQQYa/0QIQ & 	
!l!+	
 	
 $$..	'$O1t' 	dAlG-  #,A,%"719N!O Qi	
 q5 $A, ( "2 "10 $ " ) ) a& $  &!" (?q'A%  F( F1u
$w!e)V]^_]`Ta"b!cc3	d6 mmV,t007B7	"
 "oo
 
	 	e,u Rs   ,H-rE  r&   c                     | j                  t        | j                  |             | j                  D ]  }|j	                           y)zInitialize model weights.

        Args:
            needs_reset: If True, call reset_parameters() on modules (default for after to_empty()).
                If False, skip reset_parameters() (for __init__ where modules already self-initialized).
        rD  N)applyr   _init_weightsrN  r"  )rh   rE  blys      r7   rP  zSwinTransformerV2.init_weightsS  s<     	

74--;GH;; 	$C!!#	$r9   mc                    t        |t        j                        rOt        |j                  d       |j
                  +t        j                  j                  |j
                  d       yy|rt        |d      r|j                          yyy)zInitialize weights for Linear layers.

        Args:
            m: Module to initialize.
            needs_reset: Whether to call reset_parameters() on modules.
        g{Gz?)stdNr   rg   )
r  rU   r[   r   ry   rK   ro   rp   hasattrrg   )rh   rY  rE  s      r7   rW  zSwinTransformerV2._init_weights^  sg     a#!((,vv!!!!&&!, "WQ(:;  <[r9   window_ratioc                 ^   ||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 )aT  Updates the image resolution, window size, and so the pair-wise relative positions.

        Args:
            img_size (Optional[Tuple[int, int]]): New input resolution, if None current resolution is used
            patch_size (Optional[Tuple[int, int]): New patch size, if None use current patch size
            window_size (Optional[int]): New window size, if None based on new_img_size // window_div
            window_ratio (int): divisor for calculating window size from patch grid size
            always_partition: always partition / shift windows even if feat size is < window
        N)r:   r.  r)   r   r   r  )rK  r   rL  r   	enumeraterN  r   )rh   r:   r.  r%   r]  r   rL  r   indexstagestage_scales              r7   r   z SwinTransformerV2.set_input_sizel  s    " :#9++X*+U((22I<#;I Fql!2 FGK%dkk2 	LE5s519a00K  $Q<;6	!8ST'!1 ! 	 !Gs   B*c                     t               }| j                         D ]2  \  }}t        dD cg c]  }||v  c}      s"|j                  |       4 |S c c}w )zGet parameter names that should not use weight decay.

        Returns:
            Set of parameter names to exclude from weight decay.
        )r]   rY   )setnamed_modulesr   add)rh   nodnrY  kws        r7   no_weight_decayz!SwinTransformerV2.no_weight_decay  sW     e&&( 	DAq&@AB!GAB
	 
 Bs   A
coarsec                 2    t        d|rd      S g d      S )zCreate parameter group matcher for optimizer parameter groups.

        Args:
            coarse: If True, use coarse grouping.

        Returns:
            Dictionary mapping group names to regex patterns.
        z^absolute_pos_embed|patch_embedz^layers\.(\d+)))z^layers\.(\d+).downsample)r   )z^layers\.(\d+)\.\w+\.(\d+)N)z^norm)i )stemr  )rM  )rh   rk  s     r7   group_matcherzSwinTransformerV2.group_matcher  s)     3(.$
 	
5
 	
r9   enablec                 4    | j                   D ]	  }||_         y)z}Enable or disable gradient checkpointing.

        Args:
            enable: If True, enable gradient checkpointing.
        N)rN  r  )rh   ro  ls      r7   set_grad_checkpointingz(SwinTransformerV2.set_grad_checkpointing  s      	*A#)A 	*r9   c                 .    | j                   j                  S )z_Get the classifier head.

        Returns:
            The classification head module.
        )rO  fcru   s    r7   get_classifierz SwinTransformerV2.get_classifier  s     yy||r9   c                 J    || _         | j                  j                  ||       y)zReset the classification head.

        Args:
            num_classes: Number of classes for new head.
            global_pool: Global pooling type.
        N)r0  rO  reset)rh   r0  r1  s      r7   reset_classifierz"SwinTransformerV2.reset_classifier  s     '		[1r9   r$   indicesr  
stop_earlyr=  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 )a   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:

        )NCHWzOutput shape must be NCHW.Nr   r   r(   r)   )r   rF  rN  rK  rW   r  r  r_  r  r0   r1   append)rh   r$   ry  r  rz  r=  r{  intermediatestake_indices	max_index
num_stagesstagesr  ra  x_inters                  r7   forward_intermediatesz'SwinTransformerV2.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-r9   
prune_norm
prune_headc                     t        t        | j                        |      \  }}| j                  d|dz    | _        |rt        j                         | _        |r| j                  dd       |S )z@ Prune layers not required for specified intermediates.
        Nr   r   r;  )r   rF  rN  rU   r   r  rx  )rh   ry  r  r  r  r  s         r7   prune_intermediate_layersz+SwinTransformerV2.prune_intermediate_layers  s]     #7s4;;7G"Qikk.9q=1DI!!!R(r9   c                 l    | j                  |      }| j                  |      }| j                  |      }|S )zForward pass through feature extraction layers.

        Args:
            x: Input tensor of shape (B, C, H, W).

        Returns:
            Feature tensor of shape (B, H', W', C).
        )rK  rN  r  rh   r$   s     r7   forward_featuresz"SwinTransformerV2.forward_features  s3     QKKNIIaLr9   
pre_logitsc                 N    |r| j                  |d      S | j                  |      S )a  Forward pass through classification head.

        Args:
            x: Feature tensor of shape (B, H, W, C).
            pre_logits: If True, return features before final linear layer.

        Returns:
            Logits tensor of shape (B, num_classes) or pre-logits.
        T)r  )rO  )rh   r$   r  s      r7   forward_headzSwinTransformerV2.forward_head  s&     1;tyyty,L		!Lr9   c                 J    | j                  |      }| j                  |      }|S )zForward pass through the model.

        Args:
            x: Input tensor of shape (B, C, H, W).

        Returns:
            Logits tensor of shape (B, num_classes).
        )r  r  r  s     r7   r   zSwinTransformerV2.forward%  s)     !!!$a r9   )T)NNNr|   NFr   )NFFr}  F)r   FT)(r   r   r   r   rU   r   r   r   r#  r
   r   r   r   r   r   r   rT   rP  rW  r   r   rW   r  ignorer	   rj  r   r   rn  rr  ru  rx  r   r   r  r  r  r  r   r   r   s   @r7   r#   r#     s    +.#$&2)7-.%*$(!!!$&$&$'.4*,,,7C/x-'x- x- 	x-
 x- x- x- #s(Ox- S#Xx- +x- #x- "x- x- x- x-  "!x-" "#x-$ "%x-& S(]+'x-( RYY)x-* &+38_+x-t	$ 	$ 	$!ryy !t !t !  374859*+/3uS#X/ !sCx1 "%S/2	
 #3- 'tn@ YY
S 
 
 YY
D 
T#s(^ 
 
$ YY*T *T * * YY		  2C 2hsm 2W[ 2 8<$$',0 ||0  eCcN340  	0 
 0  0  !%0  
tELL!5tELL7I)I#JJ	K0 h ./$#	3S	>*  	 %,, 5<< 
Mell 
M 
M 
M %,, r9   
state_dictmodelc                     | j                  d|       } | j                  d|       } d| v }i }ddl}| j                         D ]  \  }}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      }|s&|j                  dd |      }|j                  dd      }|||<    |S c c}w )aM  Filter and process checkpoint state dict for loading.

    Handles resizing of patch embeddings and relative position tables
    when model size differs from checkpoint.

    Args:
        state_dict: Checkpoint state dictionary.
        model: Target model to load weights into.

    Returns:
        Filtered state dictionary.
    r  r  zhead.fc.weightr   N)rR   rQ   r   zpatch_embed.proj.weightr   r,   bicubicT)interpolation	antialiasverbosezlayers.(\d+).downsamplec                 D    dt        | j                  d            dz    dS )Nr?  r   z.downsample)r   group)r$   s    r7   <lambda>z&checkpoint_filter_fn.<locals>.<lambda>V  s$    ws177ST:YZGZF[[f=g r9   zhead.zhead.fc.)getreitemsr   rK  rc   ry   r.   r   subreplace)r  r  native_checkpointout_dictr  r   r   rh  r   r3   r4   s              r7   checkpoint_filter_fnr  3  s    4Jj9J(J6H  " 1 ab1Qbc$)**//66<<JAq!Qwwr{a1772;!#3(F"+"  !13gijkA		':.A'* O) cs   C;
variant
pretrainedc           	          t        d t        |j                  dd            D              }|j                  d|      }t	        t
        | |ft        t        d|      d|}|S )zCreate a Swin Transformer V2 model.

    Args:
        variant: Model variant name.
        pretrained: If True, load pretrained weights.
        **kwargs: Additional model arguments.

    Returns:
        SwinTransformerV2 model instance.
    c              3   &   K   | ]	  \  }}|  y wr   r   )r  r  r   s      r7   r  z._create_swin_transformer_v2.<locals>.<genexpr>h  s     \da\r  r3  )r   r   r   r   out_indicesT)flatten_sequentialr  )pretrained_filter_fnfeature_cfg)r   r_  r  popr   r#   r  rM  )r  r  rQ  default_out_indicesr  r  s         r7   _create_swin_transformer_v2r  ]  sj      \i

8\8Z.[\\**],?@K 7J1DkJ 	E
 Lr9   c                 4    | ddddddt         t        ddd	d
|S )Nr&  )r(      r  )r|   r|   g?r  Tzpatch_embed.projzhead.fcmit)urlr0  
input_size	pool_sizecrop_pctr  fixed_input_sizemeanr[  
first_conv
classifierlicenser   )r  rQ  s     r7   _cfgr  s  s5    =v%.B(	 # r9   ztimm/z{https://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_base_patch4_window12to16_192to256_22kto1k_ft.pth)	hf_hub_idr  z{https://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_base_patch4_window12to24_192to384_22kto1k_ft.pth)r(     r  )r,  r,  r}   )r  r  r  r  r  z|https://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_large_patch4_window12to16_192to256_22kto1k_ft.pthz|https://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_large_patch4_window12to24_192to384_22kto1k_ft.pthzfhttps://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_tiny_patch4_window8_256.pthzghttps://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_tiny_patch4_window16_256.pthzghttps://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_small_patch4_window8_256.pthzhhttps://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_small_patch4_window16_256.pthzfhttps://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_base_patch4_window8_256.pthzghttps://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_base_patch4_window16_256.pthzkhttps://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_base_patch4_window12_192_22k.pthiQU  )r(      r  )r*  r*  )r  r  r0  r  r  zlhttps://github.com/SwinTransformer/storage/releases/download/v2.0.0/swinv2_large_patch4_window12_192_22k.pth)2swinv2_base_window12to16_192to256.ms_in22k_ft_in1k2swinv2_base_window12to24_192to384.ms_in22k_ft_in1k3swinv2_large_window12to16_192to256.ms_in22k_ft_in1k3swinv2_large_window12to24_192to384.ms_in22k_ft_in1kzswinv2_tiny_window8_256.ms_in1kz swinv2_tiny_window16_256.ms_in1kz swinv2_small_window8_256.ms_in1kz!swinv2_small_window16_256.ms_in1kzswinv2_base_window8_256.ms_in1kz swinv2_base_window16_256.ms_in1k!swinv2_base_window12_192.ms_in22k"swinv2_large_window12_192.ms_in22kc           	      P    t        dddd      }t        	 dd| it        |fi |S )z"Swin-T V2 @ 256x256, window 16x16.r   r(  r)  r+  r%   r2  r3  r@   r  )swinv2_tiny_window16_256rM  r  r  rQ  
model_argss      r7   r  r    sD     "<SabJ&"Y/9Y=A*=WPV=WY Yr9   c           	      P    t        dddd      }t        	 dd| it        |fi |S )z Swin-T V2 @ 256x256, window 8x8.r|   r(  r)  r+  r  r  )swinv2_tiny_window8_256r  r  s      r7   r  r    sC     !r,R`aJ&!X.8X<@<Vv<VX Xr9   c           	      P    t        dddd      }t        	 dd| it        |fi |S )z"Swin-S V2 @ 256x256, window 16x16.r   r(  r)   r)      r)   r+  r  r  )swinv2_small_window16_256r  r  s      r7   r  r    sD     "=TbcJ&#Z0:Z>B:>XQW>XZ Zr9   c           	      P    t        dddd      }t        	 dd| it        |fi |S )z Swin-S V2 @ 256x256, window 8x8.r|   r(  r  r+  r  r  )swinv2_small_window8_256r  r  s      r7   r  r    sD     !r-SabJ&"Y/9Y=A*=WPV=WY Yr9   c           	      P    t        dddd      }t        	 dd| it        |fi |S )z"Swin-B V2 @ 256x256, window 16x16.r      r  r*   r|   r       r  r  )swinv2_base_window16_256r  r  s      r7   r  r    D     "MUcdJ&"Y/9Y=A*=WPV=WY Yr9   c           	      P    t        dddd      }t        	 dd| it        |fi |S )z Swin-B V2 @ 256x256, window 8x8.r|   r  r  r  r  r  )swinv2_base_window8_256r  r  s      r7   r  r    sC     !s=TbcJ&!X.8X<@<Vv<VX Xr9   c           	      P    t        dddd      }t        	 dd| it        |fi |S )z"Swin-B V2 @ 192x192, window 12x12.r,  r  r  r  r  r  )swinv2_base_window12_192r  r  s      r7   r  r    r  r9   c           	      R    t        ddddd      }t        	 dd| it        |fi |S )	zQSwin-B V2 @ 192x192, trained at window 12x12, fine-tuned to 256x256 window 16x16.r   r  r  r  r,  r,  r,  r*  r%   r2  r3  r@   r9  r  )!swinv2_base_window12to16_192to256r  r  s      r7   r  r    K     #m~ /1J '+b8BbFJ:F`Y_F`b br9   c           	      R    t        ddddd      }t        	 dd| it        |fi |S )	zQSwin-B V2 @ 192x192, trained at window 12x12, fine-tuned to 384x384 window 24x24.r-  r  r  r  r  r  r  )!swinv2_base_window12to24_192to384r  r  s      r7   r  r    r  r9   c           	      P    t        dddd      }t        	 dd| it        |fi |S )z"Swin-L V2 @ 192x192, window 12x12.r,  r  r  r*  r,  r-  0   r  r  )swinv2_large_window12_192r  r  s      r7   r  r    sD     "MUdeJ&#Z0:Z>B:>XQW>XZ Zr9   c           	      R    t        ddddd      }t        	 dd| it        |fi |S )	zQSwin-L V2 @ 192x192, trained at window 12x12, fine-tuned to 256x256 window 16x16.r   r  r  r  r  r  r  )"swinv2_large_window12to16_192to256r  r  s      r7   r  r    K     #m /1J ',c9CcGKJGaZ`Gac cr9   c           	      R    t        ddddd      }t        	 dd| it        |fi |S )	zQSwin-L V2 @ 192x192, trained at window 12x12, fine-tuned to 384x384 window 24x24.r-  r  r  r  r  r  r  )"swinv2_large_window12to24_192to384r  r  s      r7   r  r    r  r9   r  r  r  r  r  r  )swinv2_base_window12_192_22k)swinv2_base_window12to16_192to256_22kft1k)swinv2_base_window12to24_192to384_22kft1kswinv2_large_window12_192_22k*swinv2_large_window12to16_192to256_22kft1k*swinv2_large_window12to24_192to384_22kft1kr  )r;  )Nr   rq   	functoolsr   typingr   r   r   r   r   r	   r
   r   r   rW   torch.nnrU   torch.nn.functionalr   r   	timm.datar   r   timm.layersr   r   r   r   r   r   r   r   r   r   r   _builderr   	_featuresr   _features_fxr   _manipulater   	_registryr    r!   r"   __all__r   r   r   r8   r<   r   r>   r   r   r  r#   r#  r  r   r  r  default_cfgsr  r  r  r  r  r  r  r  r  r  r  r  r   r   r9   r7   <module>r     s     O O O     A; ; ; ; * + 3 # Y Y
#uS#X./ <<38_ \\& 38_ S/ \\	 ,Qbii QhpRYY pf&299 &RM3RYY M3`a		 aH'T#u||*;%< 'RYY 'SWX[]b]i]iXiSj 'T $ Uf , %:> J; ;? J Hs;
 <@ K< <@ K Hs< (,t( )-u) )-u) *.v* (,t( )-u)
 *.ymv*
 +/zmv+e7& 7t Y YDU Y Y X XCT X X Z$ ZEV Z Z Y YDU Y Y Y YDU Y Y X XCT X X Y YDU Y Y b$ bM^ b b b$ bM^ b b Z$ ZEV Z Z c4 cN_ c c c4 cN_ c c H$G1e1e%I2g2g' r9   