o

    hD                     @   s@   d dl mZ ddlmZmZ ddlmZmZ G dd deZdS )    N   )
BaseModule
ConvModule)
BottleneckResLayerc                       sB   e Zd ZdZdedddedddd	f fd
d	Zdd
 Z  ZS )HRFuseScalesa  Fuse feature map of multiple scales in HRNet.

    Args:
        in_channels (list[int]): The input channels of all scales.
        out_channels (int): The channels of fused feature map.
            Defaults to 2048.
        norm_cfg (dict): dictionary to construct norm layers.
            Defaults to ``dict(type='BN', momentum=0.1)``.
        init_cfg (dict | list[dict], optional): Initialization config dict.
            Defaults to ``dict(type='Normal', layer='Linear', std=0.01))``.
    i   BNg?)typemomentumNormalLinearg{Gz?)r	   layerstdc           	         s   t t| j|d || _|| _|| _t}g d}g }tt|D ]}|	t
||| || ddd q t|| _
g }tt|d D ]}|	t|| ||d  ddd| jdd qCt|| _t|d | jd| jdd	| _d S )
N)init_cfg)      i   i      )in_channelsout_channels
num_blocksstride   r   F)r   r   kernel_sizer   paddingnorm_cfgbias)r   r   r   r   r   )superr   __init__r   r   r   r   rangelenappendr   nn
ModuleListincrease_layersr   downsample_layersfinal_layer)	selfr   r   r   r   
block_typer#   ir$   	__class__ 0/root/Awesome-Backbones/configs/necks/hr_fuse.pyr      sN   

zHRFuseScales.__init__c                 C   sz   t |trt|t| jksJ | jd |d }tt| jD ]}| j| || j|d  ||d   }q | |fS )Nr   r   )
isinstancetupler   r   r#   r   r$   r%   )r&   xfeatr(   r+   r+   r,   forwardG   s    zHRFuseScales.forward)__name__
__module____qualname____doc__dictr   r1   
__classcell__r+   r+   r)   r,   r      s    
3r   )	torch.nnr!   commonr   r   Zbackbones.resnetr   r   r   r+   r+   r+   r,   <module>   s   