mediatok-player / codecs /gigatoken.py
Daankular's picture
Upload folder using huggingface_hub
92076a7 verified
Raw
History Blame Contribute Delete
6.46 kB
import torch
from .base import VideoTokenizer
class GigaTokenVideoCodec(VideoTokenizer):
"""Hierarchical vision tokenizer with large vocabulary and structured token types.
Rather than flat per-frame tokens (``[512 tokens for frame 187]``),
the stream encodes structured changes in a token hierarchy:
::
Layer 0 — Scene 4 tokens composition, lighting, environment
Layer 1 — Camera 8 tokens camera params, motion, cut boundaries
Layer 2 — Object 16 tokens object identities, positions, categories
Layer 3 — Motion 32 tokens temporal dynamics, optical flow
Layer 4 — Texture 128 tokens fine details, edges, surface patterns
Layer 5 — Residual 256 tokens reconstruction error from coarse layers
Scene and Camera tokens change rarely across consecutive frames, dramatically
reducing temporal redundancy compared to frame-by-frame encoding.
Each token is drawn from a vocabulary of ``vocab_size`` entries (default 262144).
Layers are independently decodable — a layer mask selects which layers to
reconstruct, enabling progressive quality scaling, semantic seeking, and
object-level editing directly in the compressed domain.
Backends:
"research" — random tokens matching the hierarchical layout (default)
"magvit2" — Open-MAGVIT2 262k-codebook visual tokenizer (XPU/CUDA)
"cosmos" — NVIDIA Cosmos Tokenizer (XPU/CUDA)
"""
LAYER_NAMES = ["Scene", "Camera", "Object", "Motion", "Texture", "Residual"]
def __init__(self, device: str = "cpu", vocab_size: int = 262144,
tokens_per_layer: list | None = None,
backend: str = "research"):
self.device = torch.device(device)
self._vocab_size = vocab_size
self._layers = 6
self._tpl = tokens_per_layer or [4, 8, 16, 32, 128, 256]
self._backend = backend
self._real_backend = None
self._loaded = False
def _lazy_load(self):
if self._loaded:
return
if self._backend == "magvit2":
self._real_backend = _Magvit2Backend(self.device)
elif self._backend == "cosmos":
self._real_backend = _CosmosBackend(self.device)
self._loaded = True
def encode(self, video: torch.Tensor) -> list[torch.Tensor]:
self._lazy_load()
if self._real_backend is not None:
return self._real_backend.encode(video)
B, C, T, H, W = video.shape
tokens = []
for layer in range(self._layers):
n = self._tpl[layer]
t = torch.randint(0, self._vocab_size, (B, T, n),
dtype=torch.int32, device=self.device)
tokens.append(t)
return tokens
def decode(self, tokens: list[torch.Tensor]) -> torch.Tensor:
self._lazy_load()
if self._real_backend is not None:
return self._real_backend.decode(tokens)
if not tokens:
B, T = 1, 0
elif tokens[0].dim() == 3:
B, T = tokens[0].shape[0], tokens[0].shape[1]
else:
B, T = 1, 1
H, W = 64, 64
out = torch.randn(B, 3, T, H, W, device=self.device)
return out
@property
def num_layers(self) -> int:
return self._layers
@property
def layer_token_counts(self) -> list[int]:
return self._tpl
@property
def vocab_size(self) -> int:
return self._vocab_size
@property
def name(self) -> str:
return f"gigatoken-{self._backend}-v{self._vocab_size}"
class _CosmosBackend:
def __init__(self, device):
self.device = device
self._encoder = None
self._decoder = None
self._loaded = False
def _lazy_load(self):
if self._loaded:
return
from cosmos_tokenizer.video_lib import CausalVideoTokenizer
variant = "DV8x16x16"
ckpt_enc = f"pretrained_ckpts/Cosmos-0.1-Tokenizer-{variant}/encoder.jit"
ckpt_dec = f"pretrained_ckpts/Cosmos-0.1-Tokenizer-{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
@torch.inference_mode()
def encode(self, video):
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):
self._lazy_load()
reconstructed = self._decoder.decode(tokens[0].to(self.device))
return reconstructed
class _Magvit2Backend:
"""Open-MAGVIT2 visual tokenizer with 262k LFQ codebook.
Uses Lookup-Free Quantization to produce discrete visual tokens
directly — no continuous latents. The 262144-codebook variant
is competitive with next-generation codecs in human evaluations.
Pretrained models: ``TencentARC/Open-MAGVIT2-Tokenizer-262144-Video``
"""
def __init__(self, device, variant: str = "262144"):
self.device = device
self._variant = variant
self._model = None
self._loaded = False
def _lazy_load(self):
if self._loaded:
return
try:
from open_magvit2 import get_tokenizer
repo = f"TencentARC/Open-MAGVIT2-Tokenizer-{self._variant}-Video"
self._model = get_tokenizer(repo, device=str(self.device))
self._model.eval()
self._loaded = True
except ImportError:
raise ImportError(
"open_magvit2 not installed; try: pip install open-magvit2"
)
except Exception as e:
raise RuntimeError(f"failed to load Open-MAGVIT2: {e}")
@torch.inference_mode()
def encode(self, video):
self._lazy_load()
video = video.to(self.device)
tokens = self._model.encode(video)
if isinstance(tokens, (list, tuple)):
tokens = tokens[0]
return [tokens]
@torch.inference_mode()
def decode(self, tokens):
self._lazy_load()
recon = self._model.decode(tokens[0].to(self.device))
return recon