MCPcopy Create free account
hub / github.com/Vchitect/Latte / TemporalDifferenceEncoder

Class TemporalDifferenceEncoder

tools/utils/layers.py:255–297  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

253#----------------------------------------------------------------------------
254
255class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected