MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / forward

Method forward

preprocess/auxiliary/AutoShot.py:272–306  ·  view source on GitHub ↗
(self, inputs)

Source from the content-addressed store, hash-verified

270 self.device = "cpu"
271
272 def forward(self, inputs):
273 # pt version [BS, C, N, H, W], so [3, 4] means apply avg on spatial dim, out dim is [BS, C, N]
274 x = torch.cat([torch.mean(x, dim=[3, 4]) for x in inputs], dim=1)
275
276 if self.stop_gradient:
277 x = x.detach()
278
279 x = x.permute(dims=[0, 2, 1]) # out is [BS, N ,C]
280 batch_size, time_window, old_channels = x.shape
281 x = x.reshape(shape=[batch_size * time_window, old_channels]) # [BS X N, C]
282 x = self.projection(x)
283 x = F.normalize(x, p=2, dim=1) # norm at C dim
284
285 _, new_channels = x.shape
286 x = x.reshape(shape=[batch_size, time_window, new_channels])
287 y = x.permute(dims=[0, 2, 1])
288 similarities = torch.matmul(x, y) # [batch_size, time_window, time_window]
289 # note that it operates on dimensions of the input tensor in a backward fashion (from last dimension to the first dimension)
290 similarities_padded = F.pad(similarities,
291 pad=[(self.lookup_window - 1) // 2, (self.lookup_window - 1) // 2, 0, 0, 0, 0])
292
293 batch_indices = torch.arange(0, batch_size, device=self.device). \
294 reshape(shape=[batch_size, 1, 1]). \
295 repeat([1, time_window, self.lookup_window])
296 time_indices = torch.arange(0, time_window, device=self.device). \
297 reshape(shape=[1, time_window, 1]). \
298 repeat([batch_size, 1, self.lookup_window])
299 lookup_indices = torch.arange(0, self.lookup_window, device=self.device). \
300 reshape(shape=[1, 1, self.lookup_window]). \
301 repeat([batch_size, time_window, 1]) + time_indices
302
303 indices = torch.stack([batch_indices, time_indices, lookup_indices], dim=-1)
304
305 similarities = gather_nd(similarities_padded, indices)
306 return self.fc(similarities)
307
308
309class ColorHistograms(nn.Module):

Callers

nothing calls this directly

Calls 1

gather_ndFunction · 0.85

Tested by

no test coverage detected