
    ^j.                        d Z ddlmZmZmZmZ ddlZddlZddl	m
Z
 ddlm
c mZ  G d de
j                        Zdeeeeege
j                  f   f      dedee
j                     fd	Zdd
Z G d de
j                        Z G d de
j                        Z G d de
j                        Z G d de
j                        Z G d de
j                        Zy)z
Projector.    )CallableOptionalSequenceUnionNc                   *     e Zd ZdZd fd	Zd Z xZS )	LayerNormaH  A LayerNorm variant, popularized by Transformers, that performs point-wise mean and variance normalization over
    the channel dimension for inputs that have shape (batch_size, channels, height, width).

    https://github.com/facebookresearch/ConvNeXt/blob/d1fa8f6fef0a165b27399986cc2bdacc92777e40/models/convnext.py#L119
    c                     t         |           t        j                  t	        j
                  |            | _        t        j                  t	        j                  |            | _        || _	        |f| _
        y N)super__init__nn	Parametertorchonesweightzerosbiasepsnormalized_shape)selfr   r   	__class__s      k/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/rfdetr/models/backbone/projector.pyr   zLayerNorm.__init__   sT    ll5::.>#?@LL-=!>?	!1 3    c                     |j                  dddd      }t        j                  || j                  | j                  | j
                  | j                        }|j                  dddd      }|S )zb
        LayerNorm forward
        TODO: this is a hack to avoid overflow when using fp16
        r            )permuteF
layer_normr   r   r   r   r   xs     r   forwardzLayerNorm.forward$   sY    
 IIaAq!LLD114;;		488TIIaAq!r   )gư>__name__
__module____qualname____doc__r   r#   __classcell__r   s   @r   r   r      s    4r   r   normout_channelsreturnc                 f    | yt        | t              rt        |       dk(  rydd i|    }  | |      S )z
    Args:
        norm: Either one of BN, SyncBN, FrozenBN, GN;
            or a callable that takes a channel number and returns the normalization layer as a nn.Module.

    Returns:
        The normalization layer.
    Nr   LNc                     t        |       S r
   )r   )channelss    r   <lambda>zget_norm.<locals>.<lambda>>   s    9X#6 r   )
isinstancestrlen)r+   r,   s     r   get_normr6   /   sF     |$t9>6

 r   c                    | dk(  rt        j                  |      }|S | dk(  rt        j                  |      }|S | dv rt        j                  d|      }|S | t        j                         }|S t        dj                  |             )zGet activation.siluinplacerelu)	LeakyReLU	leakyrelulrelug?zUnsupported act type: {})r   SiLUReLUr<   IdentityAttributeErrorformat)namer:   modules      r   get_activationrF   C   s    v~) M 
) M 
4	4c73
 M	 
 M 7>>tDEEr   c                   8     e Zd ZdZ	 	 	 	 	 	 	 d fd	Zd Z xZS )ConvXzConv-bn module.c
           
      d   t         t        |           t        |t              s||f}|d   dz  |d   dz  f}
t        j                  |||||
||d      | _        |	rt        j                  |      | _	        n(|rt        d|      nt        j                  |      | _	        t        |d      | _        y )	Nr   r   r   F)kernel_sizestridepaddinggroupsdilationr   r/   Tr9   )r   rH   r   r3   tupler   Conv2dconvRMSNormbnr6   BatchNorm2drF   act)r   	in_planes
out_planeskernelrK   rM   rN   rU   r    rms_normrL   r   s              r   r   zConvX.__init__U   s     	eT#%&%(f%F!9>6!9>2II	
	 jj,DG4>htZ0BNNS]D^DG!#t4r   c                     | j                  | j                  | j                  |j                                           }|S )forward.)rU   rS   rQ   
contiguousr   r"   outs      r   r#   zConvX.forwardu   s-    hhtwwtyy89:
r   )r   r   r   r   r;   FFr$   r*   s   @r   rH   rH   R   s(     5@r   rH   c                   *     e Zd ZdZd fd	Zd Z xZS )
BottleneckzStandard bottleneck.c
           
          t         |           t        ||z        }
t        ||
|d   d|||	      | _        t        |
||d   d||||	      | _        |xr ||k(  | _        y)z1ch_in, ch_out, shortcut, groups, kernels, expand.r   r   rU   r    rY   )rM   rU   r    rY   N)r   r   intrH   cv1cv2add)r   c1c2shortcutgkerU   r    rY   c_r   s              r   r   zBottleneck.__init__~   sg    a[R1qcjS[\R1q
]ef(br   c                     | j                   r#|| j                  | j                  |            z   S | j                  | j                  |            S )z1'forward()' applies the YOLOv5 FPN to input data.)rf   re   rd   r!   s     r   r#   zBottleneck.forward   s:    ,0HHq488DHHQK((O$((488A;:OOr   )Tr   r   r         ?r8   FFr$   r*   s   @r   r`   r`   {   s    )Pr   r`   c                   *     e Zd ZdZd fd	Zd Z xZS )C2fz<Faster Implementation of CSP Bottleneck with 2 convolutions.c
           	      J   	 t         
           t        ||z         _        t	        |d j                  z  dd	       _        t	        d|z    j                  z  |d	       _        t        j                  	 fdt        |      D               _
        y)z3ch_in, ch_out, number, shortcut, groups, expansion.r   r   rb   c              3   n   K   | ],  }t        j                  j                  d d	       . yw)ro         ?)rk   rl   rU   r    rY   N)r`   c).0_rU   rj   r    rY   r   ri   s     r   	<genexpr>zC2f.__init__.<locals>.<genexpr>   s;      
 tvvtvvxfYcnvww
s   25N)r   r   rc   rv   rH   rd   re   r   
ModuleListrangem)r   rg   rh   nri   rj   rl   rU   r    rY   r   s   `   `` ```r   r   zC2f.__init__   s    R!VQZA3:X`aUdffb!h
  
 
1X
 
r   c                    t        | j                  |      j                  | j                  | j                  fd            j	                  fd| j
                  D               | j                  t        j                  d            S )z.Forward pass using split() instead of chunk().r   c              3   4   K   | ]  } |d            yw)N )rw   r|   ys     r   ry   zC2f.forward.<locals>.<genexpr>   s     *a1R5*s   )	listrd   splitrv   extendr|   re   r   cat)r   r"   r   s     @r   r#   zC2f.forward   s^    !""DFFDFF#3Q78	*466**xx		!Q((r   )r   Fr   rp   r8   FFr$   r*   s   @r   rr   rr      s    F
)r   rr   c                   f     e Zd ZdZ	 	 	 	 	 ddee   dedee   dedededed	ed
df fdZd Z	 xZ
S )MultiScaleProjectorzThis module implements MultiScaleProjector in :paper:`lwdetr`.

    It creates pyramid features built on top of the input feature map.
    in_channelsr,   scale_factors
num_blocksr    rY   survival_probforce_drop_last_n_featuresr-   Nc	                 B   t         t        |           || _        || _        || _        g }	g }
d| _        |D ]  |	j                  g        |D ]!  }g }dk(  rl|j                  t        j                  ||dz  dd      t        d|dz        t        j                         t        j                  |dz  |dz  dd      g       ndk(  r-|j                  t        j                  ||dz  dd      g       nSdk(  rnMd	k(  r!|j                  t        ||d
d|      g       n'dk(  rd| _        t        dj                              t        j                   | }|	d   j                  |       $ t        j"                  |	d         |	d<   t%        t'        fd|D                    }t)        ||||      t        d|      g}t        j                   | }|
j                  |        t        j"                  |	      | _        t        j"                  |
      | _        y)a2  
        Args:
            in_channels: Channels in each input feature map level.
            out_channels: Number of channels in the output feature maps.
            scale_factors: List of scaling factors to upsample or downsample
                the input features for creating pyramid features.
        Fg      @r   )rJ   rK   r/      g       @ru   rp   r   )r    g      ?TzUnsupported scale_factor:{}r   c              3   <   K   | ]  }|t        d       z    yw)r   N)max)rw   
in_channelscales     r   ry   z/MultiScaleProjector.__init__.<locals>.<genexpr>   s     WZZ3q%=8Ws   N)r   r   r   r   r   r   use_extra_poolappendr   r   ConvTranspose2dr6   GELUrH   NotImplementedErrorrC   
Sequentialrz   rc   sumrr   stages_samplingstages)r   r   r,   r   r   r    rY   r   r   r   r   in_dimlayersr   r   s                @r   r   zMultiScaleProjector.__init__   s   $ 	!413***D'#" 9	"E""2&% .3 C<MM..vv{PQZ[\$T6Q;7GGI..v{FaKUV_`a	 c\ MM..vv{PQZ[\ c\c\MM!&&!Q:N
 d]*.D'-.K.R.RSX.YZZ/#**62].3^ #%--0C"DOBW;WWXFFL*L|,F ]]F+FMM&!s9	"v  "}}_=mmF+r   c                 B   t        |      }| j                  dk  rd| j                  rXd| j                  z
  }t        j                  j                         }t        d|      D ]  }|||dz
  z  z  }||k  sd||   dd  nL| j                  dkD  r=t        | j                        D ]%  }t        j                  ||dz             ||dz    <   ' g }t        | j                        D ]  \  }}g }	t        | j                  |         D ]  \  }
}|	j                   |||
                ! t        |	      dkD  rt        j                  |	d      }	n|	d   }	|j                   ||	              | j                  r+|j                  t!        j"                  |d   ddd             |S )	a  
        Args:
            x: Tensor of shape (N,C,H,W). H, W must be a multiple of ``self.size_divisibility``.
        Returns:
            dict[str->Tensor]:
                mapping from feature map name to pyramid feature map tensor
                in high to low resolution order. Returned feature names follow the FPN
                convention: "p<stage>", where stage has stride = 2 ** stage e.g.,
                ["p2", "p3", ..., "p6"].
        ru   r   r   N)dimr   r   )rJ   rK   rL   )r5   r   trainingnprandomuniformr{   r   r   
zeros_like	enumerater   r   r   r   r   r   
max_pool2d)r   r"   num_featuresfinal_drop_probdrop_picritical_drop_probresultsstage	feat_fusejstage_samplings               r   r#   zMultiScaleProjector.forward  s    1v#$"4"44OYY&&(F1l+  %&/\A=M*N%O"..AaDG  ,,q04::; <#..q1q5{;AE(< !$++. 	-HAuI%.t/C/CA/F%G 7!>  !!5679~!!IIiQ7	%aL	NN5+,	- NN1<<1VWXYr   )r   FFru   r   )r%   r&   r'   r(   r   rc   floatboolr   r#   r)   r*   s   @r   r   r      s      "*+X,c]X, X,  	X,
 X, X, X, X, %(X, 
X,t%r   r   c                   &     e Zd Zd fd	Zd Z xZS )SimpleProjectorc                    t         t        |           |s/t        ||dz  dd      | _        t        |dz  |dd      | _        n*t        ||ddd      | _        t        ||ddd      | _        t        d|      | _        y )	Nr   Tr8   )r    rU   )r   r   )rX   r    rU   )r   r   r/   )r   r   r   rH   convx1convx2r6   ln)r   r   out_dimfactor_kernelr   s       r   r   zSimpleProjector.__init__+  sw    ot-/
tPDK
G&QDK4U[\DKDV\]DK4)r   c                 l    | j                  | j                  | j                  |d                     }|gS )r[   r   )r   r   r   r]   s      r   r#   zSimpleProjector.forward5  s-    ggdkk$++ad"345ur   F)r%   r&   r'   r   r#   r)   r*   s   @r   r   r   *  s    *r   r   r   )r(   typingr   r   r   r   numpyr   r   torch.nnr   torch.nn.functional
functionalr   Moduler   r4   rc   r6   rF   rH   r`   rr   r   r   r   r   r   <module>r      s     6 6     		 28E#xryy0@'A"ABC SV [cdfdmdm[n (&BII &RP P )")) ).E")) EPbii r   