IM-Animation-Motion-Encoder / motion_encoder.py
Rbaerk's picture
Release TiTok motion encoder code and checkpoint
71dfe4d verified
Raw
History Blame Contribute Delete
1.54 kB
"""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()