File size: 1,537 Bytes
71dfe4d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
"""Encoder-only assembly of the original TiTok motion encoding path.

TiTokEncoder, VectorQuantizer, HW_encoder_2 retain the original implementations.
Input to encode: preprocessed frames [N, 3, 256, 256], RGB in [0, 1].
"""
from pathlib import Path
import torch
from torch import nn
from omegaconf import OmegaConf
from safetensors.torch import load_file
from encoder_blocks import TiTokEncoder, HW_encoder_2
from quantizer import VectorQuantizer

class MotionEncoder(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.encoder = TiTokEncoder(config)
        vq = config.model.vq_model
        self.latent_tokens = nn.Parameter(torch.empty(vq.num_latent_tokens, self.encoder.width))
        self.quantize = VectorQuantizer(codebook_size=vq.codebook_size,
            token_size=vq.token_size, commitment_cost=vq.commitment_cost,
            use_l2_norm=vq.use_l2_norm)

    def encode(self, frames):
        z = self.encoder(pixel_values=frames, latent_tokens=self.latent_tokens)
        return self.quantize(z)

    def forward(self, frames):
        return self.encode(frames)

    @classmethod
    def from_pretrained(cls, directory=None, device='cpu'):
        directory = Path(directory or Path(__file__).parent)
        config = OmegaConf.load(directory / 'config.yaml')
        with torch.device('meta'):
            model = cls(config)
        model.load_state_dict(load_file(str(directory / 'motion_encoder_latest.safetensors')), strict=True, assign=True)
        return model.to(device).eval()