MCPcopy Create free account
hub / github.com/drinkingcoder/FlowFormer-Official / MemoryEncoder

Class MemoryEncoder

core/FlowFormer/LatentCostFormer/encoder.py:310–368  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

308 return x
309
310class 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

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected