| 253 | #---------------------------------------------------------------------------- |
| 254 | |
| 255 | class TemporalDifferenceEncoder(nn.Module): |
| 256 | def __init__(self, cfg: DictConfig): |
| 257 | super().__init__() |
| 258 | |
| 259 | self.cfg = cfg |
| 260 | |
| 261 | if self.cfg.sampling.num_frames_per_video > 1: |
| 262 | self.d = 256 |
| 263 | self.const_embed = nn.Embedding(self.cfg.sampling.max_num_frames, self.d) |
| 264 | self.time_encoder = FixedTimeEncoder( |
| 265 | self.cfg.sampling.max_num_frames, |
| 266 | skip_small_t_freqs=self.cfg.get('skip_small_t_freqs', 0)) |
| 267 | |
| 268 | def get_dim(self) -> int: |
| 269 | if self.cfg.sampling.num_frames_per_video == 1: |
| 270 | return 1 |
| 271 | else: |
| 272 | if self.cfg.sampling.type == 'uniform': |
| 273 | return self.d + self.time_encoder.get_dim() |
| 274 | else: |
| 275 | return (self.d + self.time_encoder.get_dim()) * (self.cfg.sampling.num_frames_per_video - 1) |
| 276 | |
| 277 | def forward(self, t: torch.Tensor) -> torch.Tensor: |
| 278 | misc.assert_shape(t, [None, self.cfg.sampling.num_frames_per_video]) |
| 279 | |
| 280 | batch_size = t.shape[0] |
| 281 | |
| 282 | if self.cfg.sampling.num_frames_per_video == 1: |
| 283 | out = torch.zeros(len(t), 1, device=t.device) |
| 284 | else: |
| 285 | if self.cfg.sampling.type == 'uniform': |
| 286 | num_diffs_to_use = 1 |
| 287 | t_diffs = t[:, 1] - t[:, 0] # [batch_size] |
| 288 | else: |
| 289 | num_diffs_to_use = self.cfg.sampling.num_frames_per_video - 1 |
| 290 | t_diffs = (t[:, 1:] - t[:, :-1]).view(-1) # [batch_size * (num_frames - 1)] |
| 291 | # Note: float => round => long is necessary when it's originally long |
| 292 | const_embs = self.const_embed(t_diffs.float().round().long()) # [batch_size * num_diffs_to_use, d] |
| 293 | fourier_embs = self.time_encoder(t_diffs.unsqueeze(1)) # [batch_size * num_diffs_to_use, num_fourier_feats] |
| 294 | out = torch.cat([const_embs, fourier_embs], dim=1) # [batch_size * num_diffs_to_use, d + num_fourier_feats] |
| 295 | out = out.view(batch_size, num_diffs_to_use, -1).view(batch_size, -1) # [batch_size, num_diffs_to_use * (d + num_fourier_feats)] |
| 296 | |
| 297 | return out |
| 298 | |
| 299 | #---------------------------------------------------------------------------- |
| 300 |
nothing calls this directly
no outgoing calls
no test coverage detected