o

    h                     @   sz   d dl mZ d dlmZ d dlm  mZ d dlZddlm	Z	 d dl
T d dlmZm
Z
 G dd deZG d	d
 d
e	ZdS )    )SequenceN   )ClsHead)*)
BaseModule
ModuleListc                       s.   e Zd Z				d fdd	Zdd Z  ZS )LinearBlock        Nc                    s   t  j|d t||| _d | _d | _d | _|d ur(t|	ddi || _|d ur9t|	ddi || _|dkrFtj
|d| _d S d S )N)init_cfgtyper   )p )super__init__nnLinearfcnormactdropoutevalpopDropout)selfin_channelsout_channelsdropout_ratenorm_cfgact_cfgr
   	__class__r
   5/root/Awesome-Backbones/configs/heads/stacked_head.pyr      s   zLinearBlock.__init__c                 C   sJ   |  |}| jd ur| |}| jd ur| |}| jd ur#| |}|S N)r   r   r   r   r   xr
   r
   r!   forward#   s   






zLinearBlock.forward)r	   NNN)__name__
__module____qualname__r   r%   
__classcell__r
   r
   r   r!   r      s    r   c                       sf   e Zd ZdZddeddfdef fdd
Zd	d
 Zdd Zd
d Z	dd Z
dddZdd Z  Z
S )StackedLinearClsHeada  Classifier head with several hidden fc layer and a output fc layer.

    Args:
        num_classes (int): Number of categories.
        in_channels (int): Number of channels in the input feature map.
        mid_channels (Sequence): Number of channels in the hidden fc layers.
        dropout_rate (float): Dropout rate after each hidden fc layer,
            except the last layer. Defaults to 0.
        norm_cfg (dict, optional): Config dict of normalization layer after
            each hidden fc layer, except the last layer. Defaults to None.
        act_cfg (dict, optional): Config dict of activation function after each
            hidden layer, except the last layer. Defaults to use "ReLU".
    r	   NReLU)r   r   c                    sz   t t| jdi | |dksJ d| d|| _|| _t|ts+J dt| || _|| _	|| _
|| _|   d S )Nr   zF`num_classes` of StackedLinearClsHead must be a positive integer, got z	 instead.zH`mid_channels` of StackedLinearClsHead should be a sequence, instead of r
   )
r   r*   r   num_classesr   
isinstancer   r   mid_channelsr   r   r   _init_layers)r   r,   r   r.   r   r   r   kwargsr   r
   r!   r   =   s$   
zStackedLinearClsHead.__init__c                 C   sj   t  | _| j}| jD ]}| jt||| j| jt	| j
d |}q
| jt| jd | jdd d d d S )N)r   r   r   r	   )r   layersr   r.   appendr   r   r   copydeepcopyr   r,   )r   r   hidden_channelsr
   r
   r!   r/   W   s,   

z!StackedLinearClsHead._init_layersc                 C   s
   | j  S r"   )r2   init_weights)r   r
   r
   r!   r7   m   s   
z!StackedLinearClsHead.init_weightsc                 C   s2   t |tr	|d }| jd d D ]}||}q|S Nr1   )r-   tupler2   )r   r$   layerr
   r
   r!   
pre_logitsp   s
   

zStackedLinearClsHead.pre_logitsc                 C   s   | j d |S r8   )r2   r#   r
   r
   r!   r   w   s   zStackedLinearClsHead.fcTFc                 C   sJ   |  |}| |}|r|durtj|ddnd}n|}|r#| |S |S )af  Inference without augmentation.

        Args:
            x (tuple[Tensor]): The input features.
                Multi-stage inputs are acceptable but only the last stage will
                be used to classify. The shape of every item should be
                ``(num_samples, in_channels)``.
            softmax (bool): Whether to softmax the classification score.
            post_process (bool): Whether to do post processing the
                inference results. It will convert the output to a list.

        Returns:
            Tensor | list: The inference results.

                - If no post processing, the output is a tensor with shape
                  ``(num_samples, num_classes)``.
                - If post processing, the output is a multi-dimentional list of
                  float and the dimensions are ``(num_samples, num_classes)``.
        Nr   )dim)r;   r   Fsoftmaxpost_process)r   r$   r>   r?   	cls_scorepredr
   r
   r!   simple_testz   s   


z StackedLinearClsHead.simple_testc                 K   s,   |  |}| |}| j||fi |}|S r"   )r;   r   loss)r   r$   gt_labelr0   r@   lossesr
   r
   r!   
forward_train   s   

z"StackedLinearClsHead.forward_train)TF)r&   r'   r(   __doc__dictfloatr   r/   r7   r;   r   rB   rF   r)   r
   r
   r   r!   r*   .   s    
"r*   )typingr   torch.nnr   Ztorch.nn.functional
functionalr=   r4   cls_headr   Zconfigs.basic.activationsZconfigs.common.base_moduler   r   r   r*   r
   r
   r
   r!   <module>   s   "