(self, x_ri_bev)
| 220 | |
| 221 | |
| 222 | def forward(self, x_ri_bev): |
| 223 | x_ri = x_ri_bev[:, 0:5, :, :] |
| 224 | x_bev = x_ri_bev[:, 5:10, :, :] |
| 225 | |
| 226 | feature_ri = self.featureExtracter_RI(x_ri) |
| 227 | feature_bev = self.featureExtracter_BEV(x_bev) |
| 228 | |
| 229 | feature_ri = feature_ri.squeeze(-1) |
| 230 | feature_bev = feature_bev.squeeze(-1) |
| 231 | feature_ri = feature_ri.permute(0, 2, 1) |
| 232 | feature_bev = feature_bev.permute(0, 2, 1) |
| 233 | feature_ri = F.normalize(feature_ri, dim=-1) |
| 234 | feature_bev = F.normalize(feature_bev, dim=-1) |
| 235 | |
| 236 | feature_ri = self.norm_1(feature_ri) |
| 237 | feature_bev = self.norm_1(feature_bev) |
| 238 | |
| 239 | feature_fuse1 = feature_bev + self.attn1(feature_bev, feature_ri, feature_ri, mask=None) |
| 240 | feature_fuse1 = self.norm_2(feature_fuse1) |
| 241 | feature_fuse1 = feature_fuse1 + self.ff1(feature_fuse1) |
| 242 | |
| 243 | feature_fuse2 = feature_ri + self.attn2(feature_ri, feature_bev, feature_bev, mask=None) |
| 244 | feature_fuse2 = self.norm_3(feature_fuse2) |
| 245 | feature_fuse2 = feature_fuse2 + self.ff2(feature_fuse2) |
| 246 | |
| 247 | feature_fuse1_ext = feature_fuse1 + self.attn1_ext(feature_fuse1, feature_ri, feature_ri, mask=None) |
| 248 | feature_fuse1_ext = self.norm_2_ext(feature_fuse1_ext) |
| 249 | feature_fuse1_ext = feature_fuse1_ext + self.ff1_ext(feature_fuse1_ext) |
| 250 | |
| 251 | feature_fuse2_ext = feature_fuse2 + self.attn2_ext(feature_fuse2, feature_bev, feature_bev, mask=None) |
| 252 | feature_fuse2_ext = self.norm_3_ext(feature_fuse2_ext) |
| 253 | feature_fuse2_ext = feature_fuse2_ext + self.ff2_ext(feature_fuse2_ext) |
| 254 | |
| 255 | feature_fuse = torch.cat((feature_fuse1_ext, feature_fuse2_ext), dim=-2) |
| 256 | feature_cat_origin = torch.cat((feature_bev, feature_ri), dim=-2) |
| 257 | feature_fuse = torch.cat((feature_fuse, feature_cat_origin), dim=-1) |
| 258 | |
| 259 | feature_fuse = feature_fuse.permute(0, 2, 1) |
| 260 | |
| 261 | feature_com = feature_fuse.unsqueeze(3) |
| 262 | |
| 263 | feature_com = F.normalize(feature_com, dim=1) |
| 264 | feature_com = self.net_vlad(feature_com) |
| 265 | feature_com = F.normalize(feature_com, dim=1) |
| 266 | |
| 267 | feature_ri = feature_ri.permute(0, 2, 1) |
| 268 | feature_ri = feature_ri.unsqueeze(-1) |
| 269 | feature_ri_enhanced = self.net_vlad_ri(feature_ri) |
| 270 | feature_ri_enhanced = F.normalize(feature_ri_enhanced, dim=1) |
| 271 | |
| 272 | feature_bev = feature_bev.permute(0, 2, 1) |
| 273 | feature_bev = feature_bev.unsqueeze(-1) |
| 274 | feature_bev_enhanced = self.net_vlad_ri(feature_bev) |
| 275 | feature_bev_enhanced = F.normalize(feature_bev_enhanced, dim=1) |
| 276 | feature_com = torch.cat((feature_ri_enhanced, feature_com), dim=1) |
| 277 | feature_com = torch.cat((feature_com, feature_bev_enhanced), dim=1) |
| 278 | |
| 279 | return feature_com |
nothing calls this directly
no outgoing calls
no test coverage detected