| 308 | return x |
| 309 | |
| 310 | class MemoryEncoder(nn.Module): |
| 311 | def __init__(self, cfg): |
| 312 | super(MemoryEncoder, self).__init__() |
| 313 | self.cfg = cfg |
| 314 | |
| 315 | if cfg.fnet == 'twins': |
| 316 | self.feat_encoder = twins_svt_large(pretrained=self.cfg.pretrain) |
| 317 | elif cfg.fnet == 'basicencoder': |
| 318 | self.feat_encoder = BasicEncoder(output_dim=256, norm_fn='instance') |
| 319 | else: |
| 320 | exit() |
| 321 | self.channel_convertor = nn.Conv2d(cfg.encoder_latent_dim, cfg.encoder_latent_dim, 1, padding=0, bias=False) |
| 322 | self.cost_perceiver_encoder = CostPerceiverEncoder(cfg) |
| 323 | |
| 324 | def corr(self, fmap1, fmap2): |
| 325 | |
| 326 | batch, dim, ht, wd = fmap1.shape |
| 327 | fmap1 = rearrange(fmap1, 'b (heads d) h w -> b heads (h w) d', heads=self.cfg.cost_heads_num) |
| 328 | fmap2 = rearrange(fmap2, 'b (heads d) h w -> b heads (h w) d', heads=self.cfg.cost_heads_num) |
| 329 | corr = einsum('bhid, bhjd -> bhij', fmap1, fmap2) |
| 330 | corr = corr.permute(0, 2, 1, 3).view(batch*ht*wd, self.cfg.cost_heads_num, ht, wd) |
| 331 | #corr = self.norm(self.relu(corr)) |
| 332 | corr = corr.view(batch, ht*wd, self.cfg.cost_heads_num, ht*wd).permute(0, 2, 1, 3) |
| 333 | corr = corr.view(batch, self.cfg.cost_heads_num, ht, wd, ht, wd) |
| 334 | |
| 335 | return corr |
| 336 | |
| 337 | def forward(self, img1, img2, data, context=None): |
| 338 | # The original implementation |
| 339 | # feat_s = self.feat_encoder(img1) |
| 340 | # feat_t = self.feat_encoder(img2) |
| 341 | # feat_s = self.channel_convertor(feat_s) |
| 342 | # feat_t = self.channel_convertor(feat_t) |
| 343 | |
| 344 | imgs = torch.cat([img1, img2], dim=0) |
| 345 | feats = self.feat_encoder(imgs) |
| 346 | feats = self.channel_convertor(feats) |
| 347 | B = feats.shape[0] // 2 |
| 348 | |
| 349 | feat_s = feats[:B] |
| 350 | feat_t = feats[B:] |
| 351 | |
| 352 | B, C, H, W = feat_s.shape |
| 353 | size = (H, W) |
| 354 | |
| 355 | if self.cfg.feat_cross_attn: |
| 356 | feat_s = feat_s.flatten(2).transpose(1, 2) |
| 357 | feat_t = feat_t.flatten(2).transpose(1, 2) |
| 358 | |
| 359 | for layer in self.layers: |
| 360 | feat_s, feat_t = layer(feat_s, feat_t, size) |
| 361 | |
| 362 | feat_s = feat_s.reshape(B, *size, -1).permute(0, 3, 1, 2).contiguous() |
| 363 | feat_t = feat_t.reshape(B, *size, -1).permute(0, 3, 1, 2).contiguous() |
| 364 | |
| 365 | cost_volume = self.corr(feat_s, feat_t) |
| 366 | x = self.cost_perceiver_encoder(cost_volume, data, context) |
| 367 | |