o

    he                     @   s*   d dl Z d dlmZ G dd dejZdS )    Nc                       s*   e Zd ZdZd fdd	Zdd Z  ZS )GlobalAveragePoolingaw  Global Average Pooling neck.

    Note that we use `view` to remove extra channel after pooling. We do not
    use `squeeze` as it will also remove the batch dimension when the tensor
    has a batch dimension of size 1, which can lead to unexpected errors.

    Args:
        dim (int): Dimensions of each sample channel, can be one of {1, 2, 3}.
            Default: 2
       c                    sl   t t|   |dv sJ dd d| d|dkr"td| _d S |dkr.td| _d S td| _d S )	N)   r      z&GlobalAveragePooling dim only support z, get z	 instead.r   r   )r   r   )r   r   r   )superr   __init__nnAdaptiveAvgPool1dgapAdaptiveAvgPool2dAdaptiveAvgPool3d)selfdim	__class__ ,/root/Awesome-Backbones/configs/necks/gap.pyr      s   
zGlobalAveragePooling.__init__c                    s~   t |trt fdd|D }tdd t||D }|S t |tjr; |}||dd}t|d |S t	d)Nc                    s   g | ]}  |qS r   )r
   ).0xr
   r   r   
<listcomp>   s    z0GlobalAveragePooling.forward.<locals>.<listcomp>c                 S   s"   g | ]
\}}| |d dqS )r   )viewsize)r   outr   r   r   r   r   !   s   " r   r   z+neck inputs should be tuple or torch.tensor)

isinstancetupleziptorchTensorr
   r   r   print	TypeError)r
   inputsoutsr   r   r   forward   s   

zGlobalAveragePooling.forward)r   )__name__
__module____qualname____doc__r   r$   
__classcell__r   r   r   r   r      s    r   )r   torch.nnr   Moduler   r   r   r   r   <module>   s   