mediatok-player / codecs /video.py
Daankular's picture
Upload folder using huggingface_hub
20857b0 verified
Raw
History Blame Contribute Delete
3.15 kB
import torch
from .base import VideoTokenizer
class CosmosVideoCodec(VideoTokenizer):
def __init__(self, device: str = "cpu", variant: str = "DV8x16x16"):
self.device = torch.device(device)
self._variant = variant
self._encoder = None
self._decoder = None
self._loaded = False
def _lazy_load(self):
if self._loaded:
return
try:
from cosmos_tokenizer.video_lib import CausalVideoTokenizer
ckpt_enc = f"pretrained_ckpts/Cosmos-0.1-Tokenizer-{self._variant}/encoder.jit"
ckpt_dec = f"pretrained_ckpts/Cosmos-0.1-Tokenizer-{self._variant}/decoder.jit"
self._encoder = CausalVideoTokenizer(checkpoint_enc=ckpt_enc).to(self.device)
self._decoder = CausalVideoTokenizer(checkpoint_dec=ckpt_dec).to(self.device)
self._encoder.eval()
self._decoder.eval()
self._loaded = True
except ImportError:
raise ImportError("cosmos-tokenizer not installed; run: pip install cosmos-tokenizer")
except Exception as e:
raise RuntimeError(f"failed to load Cosmos tokenizer: {e}")
@torch.inference_mode()
def encode(self, video: torch.Tensor) -> list[torch.Tensor]:
self._lazy_load()
video = video.to(self.device)
(latent,) = self._encoder.encode(video)
tokens = latent.long() if latent.dtype in (torch.float16, torch.bfloat16, torch.float32) else latent
return [tokens]
@torch.inference_mode()
def decode(self, tokens: list[torch.Tensor]) -> torch.Tensor:
self._lazy_load()
reconstructed = self._decoder.decode(tokens[0].to(self.device))
return reconstructed
@property
def num_layers(self) -> int:
return 1
@property
def layer_token_counts(self) -> list[int]:
return [256]
@property
def name(self) -> str:
return f"cosmos-{self._variant}"
class DummyVideoCodec(VideoTokenizer):
"""Synthetic codec for testing without GPU models.
Uses random noise as tokens — validates container/entropy pipeline.
"""
def __init__(self, device: str = "cpu", layers: int = 5, tokens_per_layer: list = None):
self.device = torch.device(device)
self._layers = layers
self._tpl = tokens_per_layer or [4, 16, 32, 64, 256]
def encode(self, video: torch.Tensor) -> list[torch.Tensor]:
B, C, T, H, W = video.shape
tokens = []
for layer in range(self._layers):
n = self._tpl[layer]
t = torch.randint(0, 262144, (B, T, n), dtype=torch.int32, device=self.device)
tokens.append(t)
return tokens
def decode(self, tokens: list[torch.Tensor]) -> torch.Tensor:
B = tokens[0].shape[0]
T = tokens[0].shape[1]
H, W = 64, 64
return torch.randn(B, 3, T, H, W, device=self.device)
@property
def num_layers(self) -> int:
return self._layers
@property
def layer_token_counts(self) -> list[int]:
return self._tpl
@property
def name(self) -> str:
return "dummy"