Spaces:
Paused
Paused
| 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}") | |
| 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] | |
| def decode(self, tokens: list[torch.Tensor]) -> torch.Tensor: | |
| self._lazy_load() | |
| reconstructed = self._decoder.decode(tokens[0].to(self.device)) | |
| return reconstructed | |
| def num_layers(self) -> int: | |
| return 1 | |
| def layer_token_counts(self) -> list[int]: | |
| return [256] | |
| 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) | |
| def num_layers(self) -> int: | |
| return self._layers | |
| def layer_token_counts(self) -> list[int]: | |
| return self._tpl | |
| def name(self) -> str: | |
| return "dummy" | |