File size: 3,150 Bytes
20857b0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
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"