
    ^jq                     6   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
 ddlZddlmZ ddlmc mZ ddlmZmZmZmZ ddlmZmZmZmZmZmZmZmZmZm Z m!Z!m"Z"m#Z#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/m0Z0 ddl1m2Z2m3Z3 dgZ4 ejj                  e6      Z7de8de8dejr                  dejr                  fdZ: e-e:       dejr                  dejr                  dejr                  dee8e8f   dee8e8f   dejr                  fdZ; G d dejx                        Z= G d dejx                        Z>dejr                  de8deejr                  ee8e8f   f   fdZ?	 d8dejr                  de8d ee8e8f   d!eee8e8f      dejr                  f
d"Z@ G d# dejx                        ZAd$ ZBd9d%ZC e2 eCd&d'd(eedd)d*+       eCd,d'd(eedd)d*+       eCd-d'd(eedd)d*+       eCeed.d/d01      d2      ZDd:d3ZEe3d:deAfd4       ZFe3d:deAfd5       ZGe3d:deAfd6       ZHe3d:deAfd7       ZIy);a+   Vision Transformer (ViT) in PyTorch

A PyTorch implement of Vision Transformers as described in:

'Exploring Plain Vision Transformer Backbones for Object Detection'
    - https://arxiv.org/abs/2203.16527

'Segment Anything Model (SAM)'
    - https://github.com/facebookresearch/segment-anything/

    N)partial)CallableListOptionalTupleTypeUnion)IMAGENET_DEFAULT_MEANIMAGENET_DEFAULT_STDIMAGENET_INCEPTION_MEANIMAGENET_INCEPTION_STD)
PatchEmbedMlpDropPathcalculate_drop_path_ratesPatchDropoutLayerNorm2d
LayerScaleClassifierHeadNormMlpClassifierHeadFormatresample_abs_pos_embed_nhwcRotaryEmbeddingCatapply_rot_embed_cat	to_2tupleuse_fused_attn)Final   )build_model_with_cfg)feature_take_indices)register_notrace_function)
checkpointcheckpoint_seq)generate_default_cfgsregister_modelVisionTransformerSAMq_sizek_sizerel_posreturnc                    t        dt        | |      z  dz
        }|j                  d   |k7  rjt        j                  |j                  d|j                  d   d      j                  ddd      |d      }|j                  d|      j                  dd      }n|}t        j                  | t        j                        dddf   t        || z  d	      z  }t        j                  |t        j                        dddf   t        | |z  d	      z  }||z
  |dz
  t        | |z  d	      z  z   }||j                            S )
a\  
    Get relative positional embeddings according to the relative positions of
        query and key sizes.
    Args:
        q_size (int): size of query q.
        k_size (int): size of key k.
        rel_pos (Tensor): relative position embeddings (L, C).

    Returns:
        Extracted positional embeddings according to relative positions.
       r   r   linear)sizemode)dtypeN      ?)intmaxshapeFinterpolatereshapepermutetorcharangefloat32long)r'   r(   r)   max_rel_distrel_pos_resizedq_coordsk_coordsrelative_coordss           m/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/vision_transformer_sam.pyget_rel_posrD   4   s*    q3vv..23L}}Q<'--OOAw}}Q/4<<Q1E

 *11"lCKKAqQ! ||F%--8DACQWY\D]]H||F%--8qACQWY\D]]H(*vzS&RU=V.VVO?//122    q	rel_pos_h	rel_pos_wc                 j   |\  }}|\  }}t        |||      }	t        |||      }
| j                  \  }}}| j                  ||||      }t        j                  d||	      }t        j                  d||
      }|dddddddddf   |dddddddddf   z   }|j                  d||z  ||z        S )a  
    Calculate decomposed Relative Positional Embeddings from :paper:`mvitv2`.
    https://github.com/facebookresearch/mvit/blob/19786631e330df9f3622e5402b4a419a263a2c80/mvit/models/attention.py
    Args:
        q (Tensor): query q in the attention layer with shape (B, q_h * q_w, C).
        rel_pos_h (Tensor): relative position embeddings (Lh, C) for height axis.
        rel_pos_w (Tensor): relative position embeddings (Lw, C) for width axis.
        q_size (Tuple): spatial sequence size of query q with (q_h, q_w).
        k_size (Tuple): spatial sequence size of key k with (k_h, k_w).

    Returns:
        bias (Tensor): attention bias to add to attention map
    zbhwc,hkc->bhwkzbhwc,wkc->bhwkNr-   )rD   r5   r8   r:   einsum)rF   rG   rH   r'   r(   q_hq_wk_hk_wRhRwB_dimr_qrel_hrel_w	attn_biass                     rC   get_decomposed_rel_pos_biasrX   W   s    ( HCHC	S#y	)B	S#y	)BIAq#
))AsC
%CLL)33ELL)33EaAq$&'%1aq0@*AAIRsC#I66rE   c                        e Zd ZU ee   ed<   dddddej                  dddddfdeded	ed
ede	de	de
ej                     dedeeeef      deej                     f fdZd Z xZS )	Attention
fused_attn   TF        NrS   	num_headsqkv_biasqk_norm	attn_drop	proj_drop
norm_layeruse_rel_pos
input_sizeropec                    ||d}t         |           ||z  dk(  sJ d       || _        ||z  | _        | j                  dz  | _        t               | _        t        j                  ||dz  fd|i|| _	        |r || j                  fi |nt        j                         | _        |r || j                  fi |nt        j                         | _        t        j                  |      | _        t        j                  ||fi || _        t        j                  |      | _        || _        | j"                  r|
J |	J d       t        j$                  t'        j(                  d|	d   z  d	z
  | j                  fi |      | _        t        j$                  t'        j(                  d|	d	   z  d	z
  | j                  fi |      | _        |
| _        y )
Ndevicer1   r   z$dim should be divisible by num_headsg         biaszBInput size must be provided if using relative positional encoding.r,   r   )super__init__r^   head_dimscaler   r[   nnLinearqkvIdentityq_normk_normDropoutra   projrb   rd   	Parameterr:   zerosrG   rH   rf   )selfrS   r^   r_   r`   ra   rb   rc   rd   re   rf   ri   r1   dd	__class__s                 rC   rm   zAttention.__init__|   s    /Y!#K%KK#"y(]]d*
(*99S#'??B?9@j5"5bkkm9@j5"5bkkmI.IIc3-"-	I.&<<&TST&  \\%++a*Q-6G!6KT]]*a^`*abDN\\%++a*Q-6G!6KT]]*a^`*abDN	rE   c                    |j                   \  }}}}||z  }|j                  ||d      }| j                  |      j                  ||d| j                  d      j                  ddddd      }|j                  d|| j                  z  |d      j                  d      \  }}	}
| j                  |      | j                  |	      }	}| j                  r(t        || j                  | j                  ||f||f      }n^d }| j                  P| j                  j                         }t        ||      j!                  |
      }t        |	|      j!                  |
      }	| j"                  rQt$        j&                  j(                  j+                  ||	|
|| j,                  r| j.                  j0                  nd      }nS|| j2                  z  }||	j5                  d	d      z  }|||z   }|j7                  d
      }| j/                  |      }||
z  }|j                  || j                  |d      j5                  dd      j                  ||d      }| j9                  |      }| j;                  |      }|j                  |||d      }|S )Nr-   rj   r,   r   r      r]   )	attn_mask	dropout_p)rS   )r5   r8   rr   viewr^   r9   unbindrt   ru   rd   rX   rG   rH   rf   	get_embedr   type_asr[   r:   rp   
functionalscaled_dot_product_attentiontrainingra   pro   	transposesoftmaxrw   rb   )rz   xrQ   HWrR   Nrr   rF   kvrW   rf   attns                 rC   forwardzAttention.forward   s6   WW
1aEIIaBhhqkq!Q;CCAq!QPQR++aT^^!3Q;BB1E1a{{1~t{{1~13At~~t~~XY[\W]`acd_efIIyy$yy**,'4088;'4088;??##@@1a#.2mm$..** A A DJJAq{{2r**D$i'<<B<'D>>$'DqAFF1dnna,66q!<DDQ2NIIaLNN1FF1aBrE   )__name__
__module____qualname__r   bool__annotations__rp   	LayerNormr3   floatr   Moduler   r   rm   r   __classcell__r|   s   @rC   rZ   rZ   y   s    d
 !!!!*,,, %48(,&& & 	&
 & & & RYY& & !sCx1& 299%&P&rE   rZ   c                        e Zd Zdddddddej                  ej
                  eddddddfdeded	ed
e	de	dedede
e   dedeej                     deej                     deej                     de	def fdZd Z xZS )Block      @TFr]   Nr   rS   r^   	mlp_ratior_   r`   rb   ra   init_values	drop_path	act_layerrc   	mlp_layerrd   window_sizec                 J   ||d}t         |           || _         ||fi || _        t	        |f||||||||dk(  r|n||f|d	|| _        |rt        |fd|i|nt        j                         | _	        |	dkD  rt        |	      nt        j                         | _         ||fi || _         |d|t        ||z        |
|d|| _        |rt        |fd|i|nt        j                         | _        |	dkD  rt        |	      | _        y t        j                         | _        y )Nrh   r   )	r^   r_   r`   ra   rb   rc   rd   re   rf   r   r]   )in_featureshidden_featuresr   drop )rl   rm   r   norm1rZ   r   r   rp   rs   ls1r   
drop_path1norm2r3   mlpls2
drop_path2)rz   rS   r^   r   r_   r`   rb   ra   r   r   r   rc   r   rd   r   re   rf   ri   r1   r{   r|   s                       rC   rm   zBlock.__init__   sC   * /&*r*

!#%0A%5zK;U
 
	 FQ:cA{AbAVXVaVaVc1:R(9-R[[]*r*
 
i0	

 
 FQ:cA{AbAVXVaVaVc1:R(9-R[[]rE   c           
      2   |j                   \  }}}}|}| j                  |      }d }| j                  dkD  rt        || j                        \  }}| j	                  | j                  | j                  |                  }| j                  dkD  rt        || j                  ||f|      }||z   }|j                  |||z  d      }|| j                  | j                  | j                  | j                  |                        z   }|j                  |||d      }|S )Nr   r-   )r5   r   r   window_partitionr   r   r   window_unpartitionr8   r   r   r   r   )rz   r   rQ   r   r   rR   shortcutpad_hws           rC   r   zBlock.forward  s    WW
1aJJqM,0a(D,<,<=IAvOODHHTYYq\23 a"1d&6&6AGAqLIIaQ#$**Q-)@ ABBIIaAr"rE   )r   r   r   rp   GELUr   r   r3   r   r   r   r   r   rm   r   r   r   s   @rC   r   r      s      "!!!!+/!)+*,,,), % '2S2S 2S 	2S
 2S 2S 2S 2S "%2S 2S BII2S RYY2S BII2S 2S 2ShrE   r   r   r   c           	      L   | j                   \  }}}}|||z  z
  |z  }|||z  z
  |z  }t        j                  | ddd|d|f      } ||z   ||z   }	}| j                  |||z  ||	|z  ||      } | j	                  dddddd      j                         j                  d|||      }
|
||	ffS )aU  
    Partition into non-overlapping windows with padding if needed.
    Args:
        x (tensor): input tokens with [B, H, W, C].
        window_size (int): window size.

    Returns:
        windows: windows after partition with [B * num_windows, window_size, window_size, C].
        (Hp, Wp): padded height and width before partition
    r   r   rj   r,   r~      r-   )r5   r6   padr   r9   
contiguous)r   r   rQ   r   r   Cpad_hpad_wHpWpwindowss              rC   r   r     s     JAq!Q1{?*k9E1{?*k9E	a!Q5!U+,AYE	B	q"#["2C[RSTAii1aAq)446;;B[Z[\GRHrE   r   hwr   c                 :   ||n|\  }}|\  }}| j                   d   ||z  |z  |z  z  }| j                  |||z  ||z  ||d      }	|	j                  dddddd      j                         j                  |||d      }	|	ddd|d|ddf   j                         }	|	S )	a  
    Window unpartition into original sequences and removing padding.
    Args:
        windows (tensor): input tokens with [B * num_windows, window_size, window_size, C].
        window_size (int): window size.
        pad_hw (Tuple): padded height and width (Hp, Wp).
        hw (Tuple): original height and width (H, W) before padding.

    Returns:
        x: unpartitioned sequences with [B, H, W, C].
    Nr   r-   r   rj   r,   r~   r   )r5   r   r9   r   )
r   r   r   r   r   r   r   r   rQ   r   s
             rC   r   r   3  s     )VrFBDAqaR"W3{BCAQk)2+<k;XZ[A			!Q1a#..055aRDA	!RaR!Q,""$AHrE   c            H           e Zd ZdZdddddddddd	d
d	ddddddd eeej                  d	      ej                  ej                  eedd	d	dddd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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j&                        d'eeej&                        d(eej&                     d)eej&                     d*ed+ed,ed-ed.eed/f   d0ed1ed2ee   d3eeeeef   eeef   f      fB fd4Zej.                  j0                  d5        Zej.                  j0                  dId6       Zej.                  j0                  dJd7       Zej.                  j0                  d8ej&                  fd9       ZdKded1ee   fd:Z	 	 	 	 	 dLd;ej<                  d<eeee e   f      d=ed>ed?ed@ed8ee ej<                     eej<                  e ej<                     f   f   fdAZ!	 	 	 dMd<eeee e   f      dBedCefdDZ"dE Z#dIdFefdGZ$dH Z% xZ&S )Nr&   z Vision Transformer for Segment-Anything Model(SAM)

    A PyTorch impl of : `Exploring Plain Vision Transformer Backbones for Object Detection` or `Segment Anything Model (SAM)`
        - https://arxiv.org/abs/2010.11929
          rj         r   TFNr]    )
output_fmtstrict_img_size   r      avgimg_size
patch_sizein_chansnum_classes	embed_dimdepthr^   r   r_   r`   r   pre_norm	drop_ratepos_drop_ratepatch_drop_rateproj_drop_rateattn_drop_ratedrop_path_rateweight_initembed_layerrc   r   block_fnr   use_abs_posrd   use_roper   global_attn_indexes.
neck_chansglobal_poolhead_hidden_sizeref_feat_shapec$                 <   t         +|           |"|#d}$|xs t        t        j                  d      }|xs t        j
                  }|| _        || _        || _        |x| _	        x| _
        | _        d| _         |d&||||| d|$| _        | j                  j                  }%t        | j                  d      r| j                  j!                         n|}&|r8t        j"                  t%        j&                  d|%d   |%d   |fi |$      | _        nd	| _        t        j*                  |
      | _        |dkD  rt/        |d      | _        nt        j2                         | _        |r	 ||fi |$nt        j2                         | _        |rt|rJ d       |!-t7        |!      dk(  sJ t9        |!d         }'t9        |!d         }(nd	x}'}(t;        ||z  d|%|'      | _        t;        ||z  dt9        |      |(      | _        nd	| _        d	| _        tA        ||      })t        jB                  tE        |      D *cg c]]  }* |d&i d|d|d|d|	d|
d|d|d|d|)|*   d|d|d|d|d|*|vr|ndd|%d|*|vr| j>                  n| j<                  |$_ c}* | _#        tE        |      D *cg c]  }*tI        d|* ||&        c}*| _%        |rjt        jB                  t        jL                  ||fddd!|$tO        |fi |$t        jL                  ||fd"ddd#|$tO        |fi |$      | _(        || _	        n/| rt        j2                         | _(        ntO        |fi |$| _(        |}| rtS        ||f| ||d$|$| _*        y	tW        ||f||d%|$| _*        y	c c}*w c c}*w )'a  
        Args:
            img_size: Input image size.
            patch_size: Patch size.
            in_chans: Number of image input channels.
            num_classes: Number of classes for classification head.
            global_pool: Type of global pooling for final sequence (default: 'token').
            embed_dim: Transformer embedding dimension.
            depth: Depth of transformer.
            num_heads: Number of attention heads.
            mlp_ratio: Ratio of mlp hidden dim to embedding dim.
            qkv_bias: Enable bias for qkv projections if True.
            init_values: Layer-scale init values (layer-scale enabled if not None).
            drop_rate: Head dropout rate.
            pos_drop_rate: Position embedding dropout rate.
            attn_drop_rate: Attention dropout rate.
            drop_path_rate: Stochastic depth rate.
            weight_init: Weight initialization scheme.
            embed_layer: Patch embedding layer.
            norm_layer: Normalization layer.
            act_layer: MLP activation layer.
            block_fn: Transformer block layer.
            use_abs_pos: If True, use absolute positional embeddings.
            use_rel_pos: If True, add relative positional embeddings to the attention map.
            use_rope: If True, add rotary position embeddings to q/k in attention block.
            window_size: Window size for window attention blocks. If 0, not use window attention.
            global_attn_indexes: Indexes for blocks using global attention. Used when window_size > 0.
            global_pool: Global pooling type.
            head_hidden_size: If set, use NormMlpHead
            ref_feat_shape: Tuple of reference feature shapes for ROPE, (global, local)
        rh   gư>)epsF)r   r   r   r   rk   
feat_ratior   r   N)r   )num_prefix_tokenszCROPE and relative pos embeddings should not be enabled at same timer,   )	in_pixels
feat_shaper   rS   r^   r   r_   r`   r   rb   ra   r   rc   r   r   rd   r   re   rf   zblocks.)modulenum_chs	reduction)kernel_sizerk   rj   )r   paddingrk   )hidden_size	pool_typer   )r   r   r   ),rl   rm   r   rp   r   r   r   r   r   num_featuresr   r   grad_checkpointingpatch_embed	grid_sizehasattrr   rx   r:   ry   	pos_embedrv   pos_dropr   
patch_droprs   norm_prelenr   r   rope_globalrope_windowr   
Sequentialrangeblocksdictfeature_infoConv2dr   neckr   headr   ),rz   r   r   r   r   r   r   r^   r   r_   r`   r   r   r   r   r   r   r   r   r   r   rc   r   r   r   r   rd   r   r   r   r   r   r   r   ri   r1   r{   r   rref_feat_shape_globalref_feat_shape_windowdprir|   s,                                              rC   rm   zVisionTransformerSAM.__init__Q  sm   J 	/B72<<T#B
(	& &ENNND1DN"'& 
!
 
 $$..	-4T5E5E|-TD'')Zd\\%++a1yQR|U^*ebd*efDN!DN

]3Q*"#DO
 !kkmDO7?
933R[[]"i$ii?)>*a///(1.2C(D%(1.2C(D%@DD%(=1Y&$4	 D  2Y&$[14	 D  $D#D (>mm( 5\)&#( '  # $ "	
   ( ) ) a& & $ $ ( ,-4G+GKQ %  *+2E)ET%%4K[K[#&# $, QVV[P\^KLD'!yAF^ 		 !"	
  J-"-		 !"  J-"-#DI& !+DKKM	 (	8R8	"J - -%# DI ' &#	
 DIA&#*^s   2A"N*Nc                 
    ddhS )Nr   
dist_tokenr   rz   s    rC   no_weight_decayz$VisionTransformerSAM.no_weight_decay  s    \**rE   c                      t        dddg      S )Nz^pos_embed|patch_embed)z^blocks\.(\d+)N)z^norm)i )stemr  )r  )rz   coarses     rC   group_matcherz"VisionTransformerSAM.group_matcher!  s    *-/CD
 	
rE   c                     || _         y N)r   )rz   enables     rC   set_grad_checkpointingz+VisionTransformerSAM.set_grad_checkpointing(  s
    "(rE   r*   c                     | j                   S r  r  r  s    rC   get_classifierz#VisionTransformerSAM.get_classifier,  s    yyrE   c                 J    || _         | j                  j                  ||       y r  )r   r  reset)rz   r   r   s      rC   reset_classifierz%VisionTransformerSAM.reset_classifier0  s    &		[1rE   r   indicesnorm
stop_earlyr   intermediates_onlyc                    |dk(  sJ d       g }t        t        | j                        |      \  }}	| j                  |      }| j                  &|t        | j                  |j                  dd       z   }| j                  |      }| j                  |      }| j                  |      }t        j                  j                         s|s| j                  }
n| j                  d|	dz    }
t        |
      D ]  \  }}| j                  r+t        j                  j                         st        ||      }n ||      }||v sJ|r3|j!                  | j#                  |j%                  dddd                   |j!                  |j%                  dddd              |r|S | j#                  |j%                  dddd            }||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 all 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 for ViT-SAM must be NCHW.Nr   rj   r   r,   )r    r   r  r   r   r   r5   r   r   r   r:   jitis_scripting	enumerater   r"   appendr  r9   )rz   r   r  r   r!  r   r"  intermediatestake_indices	max_indexr  r  blks                rC   forward_intermediatesz*VisionTransformerSAM.forward_intermediates4  s   * V#M%MM#"6s4;;7G"Qi Q>>%/!MMAMM!OOAMM!99!!#:[[F[[)a-0F' 	@FAs&&uyy/E/E/GsA&FL  "((199Q1a3H)IJ!((1aA)>?	@   IIaii1a+,-rE   
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    r   r  rp   rs   r  r  )rz   r  r.  r/  r*  r+  s         rC   prune_intermediate_layersz.VisionTransformerSAM.prune_intermediate_layerso  s]     #7s4;;7G"Qikk.9q=1DI!!!R(rE   c                    | j                  |      }| j                  &|t        | j                  |j                  dd       z   }| j	                  |      }| j                  |      }| j                  |      }| j                  r5t        j                  j                         st        | j                  |      }n| j                  |      }| j                  |j                  dddd            }|S )Nr   rj   r   r,   )r   r   r   r5   r   r   r   r   r:   r%  r&  r#   r  r  r9   rz   r   s     rC   forward_featuresz%VisionTransformerSAM.forward_features  s    Q>>%/!MMAMM!OOAMM!""599+A+A+Ct{{A.AAAIIaii1a+,rE   
pre_logitsc                 N    |r| j                  |d      S | j                  |      S )NT)r5  r  )rz   r   r5  s      rC   forward_headz!VisionTransformerSAM.forward_head  s$    0:tyyty,L		!LrE   c                 J    | j                  |      }| j                  |      }|S r  )r4  r7  r3  s     rC   r   zVisionTransformerSAM.forward  s'    !!!$a rE   F)Tr  )NFFr$  F)NFT)'r   r   r   __doc__r   r   r   NHWCrp   r   r   r   r   r3   r   r   r   strr   r   r   rm   r:   r%  ignorer  r  r  r  r  Tensorr	   r   r-  r1  r4  r7  r   r   r   s   @rC   r&   r&   J  s    ! " !!!+/"!#%%'$&$&$&!+2:&++gl+m46LL3577(-), $ %"!35!$.2PTIJJ J 	J
 J J J J J J J "%J J J !J  #!J" "#J$ "%J& "'J( )J* bii+J, !bii1-J.  RYY0/J0 299o1J2 BII3J4 5J6 7J8 9J: ;J< "'sCx=J> ?J@ AJB 'smCJD %U5c?E#s(O+K%LMEJX YY+ + YY
 
 YY) ) YY		  2C 2hsm 2 8<$$',9 ||9  eCcN349  	9 
 9  9  !%9  
tELL!5tELL7I)I#JJ	K9 z 8<$#	eCcN34  	"M$ MrE   c                     d| v }i }| j                         D ]6  \  }}|j                  d      r|dd }|j                  dd      }n|r2|||<   8 |S )z Remap SAM checkpoints -> timm z%image_encoder.patch_embed.proj.weightzimage_encoder.r   Nzmlp.linzmlp.fc)items
startswithreplace)
state_dictmodelsam_checkpointout_dictr   r   s         rC   checkpoint_filter_fnrG    sm    
 =
JNH  " 1<<()"#A		)X.A OrE   c                 2    | ddd dddt         t        ddd|S )	N  rj   r   r   ?bicubicTzpatch_embed.projzhead.fc)urlr   re   	pool_sizecrop_pctinterpolationfixed_input_sizemeanstd
first_conv
classifier)r   r   )rM  kwargss     rC   _cfgrW    s2    ?'0F(	  rE   zDhttps://dl.fbaipublicfiles.com/segment_anything/sam_vit_b_01ec64.pthztimm/z
apache-2.0rJ  r2   )rM  	hf_hub_idlicenserR  rS  r   re   rO  zDhttps://dl.fbaipublicfiles.com/segment_anything/sam_vit_l_0b3195.pthzDhttps://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pthrI  )rj      rZ  rK  )rR  rS  r   re   rO  )zsamvit_base_patch16.sa1bzsamvit_large_patch16.sa1bzsamvit_huge_patch16.sa1bsamvit_base_patch16_224c                 n    |j                  dd      }t        t        | |ft        t	        |d      d|S )Nout_indicesrj   getter)r]  feature_cls)pretrained_filter_fnfeature_cfg)popr   r&   rG  r  )variant
pretrainedrV  r]  s       rC   _create_vision_transformerre    sF    **]A.K 2[hG  rE   c           
      `    t        ddddg dddd      }t        	 d
d	| it        |fi |}|S )z# ViT-B/16 for Segment-Anything
    r   r   r   r,   r   r\      r   Tr   r   r   r   r^   r   r   rd   r   rd  )samvit_base_patch16r  re  rd  rV  
model_argsrD  s       rC   rj  rj    sR     B"R_D4J 'T*4T8<Z8R68RTELrE   c           
      `    t        ddddg dddd      }t        	 d	d| it        |fi |}|S )
z# ViT-L/16 for Segment-Anything
    r   r      )r   rh        r   Tri  rd  )samvit_large_patch16rk  rl  s       rC   rr  rr    sR     R2SbD4J 'U+5U9=j9SF9SUELrE   c           
      `    t        ddddg dddd      }t        	 d
d	| it        |fi |}|S )z# ViT-H/16 for Segment-Anything
    r   i       )      rq     r   Tr   ri  rd  )samvit_huge_patch16rk  rl  s       rC   rx  rx    sR     R2SbD4J 'T*4T8<Z8R68RTELrE   c                 d    t        ddddg dddddd	

      }t        	 dd| it        |fi |}|S )z# ViT-B/16 based on samvit arch
    r   r   r   rg  r   TFrZ  N)
r   r   r   r^   r   r   rd   r   r   r   rd  )r[  rk  rl  s       rC   r[  r[    sW     B"R_DecVZJ '!X.8X<@<Vv<VXELrE   r  )r   r9  )Jr:  logging	functoolsr   typingr   r   r   r   r   r	   r:   torch.nnrp   torch.nn.functionalr   r6   	timm.datar
   r   r   r   timm.layersr   r   r   r   r   r   r   r   r   r   r   r   r   r   r   	torch.jitr   _builderr   	_featuresr    _features_fxr!   _manipulater"   r#   	_registryr$   r%   __all__	getLoggerr   _loggerr3   r>  rD   rX   r   rZ   r   r   r   r&   rG  rW  default_cfgsre  rj  rr  rx  r[  r   rE   rC   <module>r     s  
   ? ?     r r    "  * + 3 3 < "
" '

H
%3 3S 35<< 3ELL 3@ + &7<<7<<7 <<7 c3h	7
 c3h7 \\7DQ		 QhLBII L^ 3 5uUXZ]U]A^;_ 0 gk\\(+16sCxJRSXY\^aYaSbJc
\\.K299 K\
$ % !%R"(<!"S!2 "&R"(<!"S"2 !%R"(<!"S!2  $"(<$ 3 0-& 8	 	7K 	 	 	8L 	 	 	7K 	 	 	;O 	 	rE   