o

    h                     @   sv   d dl Z d dlmZ d dlmZ d dlm  mZ d dlm	Z	 d dl
mZ ddlm
Z
 ddlmZ G d	d
 d
eZdS )    N)OrderedDict)build_activation_layer)
trunc_normal_   )
Sequential   )ClsHeadc                       sf   e Zd ZdZdeddeddddf fd	d
	Zdd Z fd
dZdd ZdddZ	dd Z
  ZS )VisionTransformerClsHeada  Vision Transformer classifier head.

    Args:
        num_classes (int): Number of categories excluding the background
            category.
        in_channels (int): Number of channels in the input feature map.
        hidden_dim (int): Number of the dimensions for hidden layer. Only
            available during pre-training. Default None.
        act_cfg (dict): The activation config. Only available during
            pre-training. Defaults to Tanh.
    NTanh)typeConstantLinearr   )r   layervalc                    sX   t t| j|d|i| || _|| _|| _|| _| jdkr&td| d|   d S )Ninit_cfgr   znum_classes=z must be a positive integer)	superr	   __init__in_channelsnum_classes
hidden_dimact_cfg
ValueError_init_layers)selfr   r   r   r   r   argskwargs	__class__ @/root/Awesome-Backbones/configs/heads/vision_transformer_head.pyr      s    


z!VisionTransformerClsHead.__init__c                 C   sh   | j d u rdt| j| jfg}ndt| j| j fdt| jfdt| j | jfg}tt|| _	d S )Nhead
pre_logitsact)
r   nnr
   r   r   r   r   r   r   layers)r   r$   r   r   r   r   0   s   
z%VisionTransformerClsHead._init_layersc                    sV   t t|   t| jdr)t| jjjt	d| jjj
 d tj
| jjj d S d S )Nr!   r   )std)r   r	   init_weightshasattrr$   r   r!   weightmathsqrtin_featuresr#   initzeros_bias)r   r   r   r   r&   ;   s   z%VisionTransformerClsHead.init_weightsc                 C   s@   t |tr	|d }|\}}| jd u r|S | j|}| j|S )N)
isinstancetupler   r$   r!   r"   )r   x_	cls_tokenr   r   r   r!   E   s   

z#VisionTransformerClsHead.pre_logitsTFc                 C   sL   |  |}| j|}|r|durtj|ddnd}n|}|r$| |S |S )a  Inference without augmentation.

        Args:
            x (tuple[tuple[tensor, tensor]]): The input features.
                Multi-stage inputs are acceptable but only the last stage will
                be used to classify. Every item should be a tuple which
                includes patch token and cls token. The cls token will be used
                to classify and the shape of it 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$   r    Fsoftmaxpost_process)r   r2   r7   r8   	cls_scorepredr   r   r   simple_testO   s   

z$VisionTransformerClsHead.simple_testc                 K   s.   |  |}| j|}| j||fi |}|S )N)r!   r$   r    loss)r   r2   gt_labelr   r9   lossesr   r   r   
forward_trains   s   
z&VisionTransformerClsHead.forward_train)TF)__name__
__module____qualname____doc__dictr   r   r&   r!   r;   r?   
__classcell__r   r   r   r   r	      s    


$r	   )r)   collectionsr   torch.nnr#   Ztorch.nn.functional
functionalr6   configs.basic.build_layerr   Zcore.initialize.weight_initr   Zcommon.base_moduler   cls_headr   r	   r   r   r   r   <module>   s   