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"