(self, inputs)
| 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 | |
| 309 | class ColorHistograms(nn.Module): |
nothing calls this directly
no test coverage detected