| 151 | """ |
| 152 | |
| 153 | def __init__( |
| 154 | self, |
| 155 | in_channels, |
| 156 | out_channels, |
| 157 | scale_factor=4, |
| 158 | norm_cfg=dict(type="BN"), |
| 159 | act_cfg=dict(type="ReLU"), |
| 160 | init_cfg=None, |
| 161 | ): |
| 162 | super(FeatureFusionModule, self).__init__(init_cfg=init_cfg) |
| 163 | channels = out_channels // scale_factor |
| 164 | self.conv0 = ConvModule(in_channels, out_channels, 1, norm_cfg=norm_cfg, act_cfg=act_cfg) |
| 165 | self.attention = nn.Sequential( |
| 166 | nn.AdaptiveAvgPool2d((1, 1)), |
| 167 | ConvModule(out_channels, channels, 1, norm_cfg=None, bias=False, act_cfg=act_cfg), |
| 168 | ConvModule(channels, out_channels, 1, norm_cfg=None, bias=False, act_cfg=None), |
| 169 | nn.Sigmoid(), |
| 170 | ) |
| 171 | |
| 172 | def forward(self, spatial_inputs, context_inputs): |
| 173 | inputs = torch.cat([spatial_inputs, context_inputs], dim=1) |