MCPcopy Create free account
hub / github.com/BIT-MJY/CVTNet / forward

Method forward

modules/cvtnet.py:222–279  ·  view source on GitHub ↗
(self, x_ri_bev)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected